From 72bd5cb152ca1df07f14a14e14a2816e188874a8 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Mon, 13 Apr 2026 11:11:08 -0700 Subject: [PATCH] feat: add compression_interception callback for LiteLLM Proxy Add a proxy callback that automatically compresses incoming /v1/messages payloads above a configurable token threshold, runs the retrieval tool loop server-side, and returns the final response. This brings compress() support to proxy deployments (e.g. Claude Code via /v1/messages). - New callback: litellm/integrations/compression_interception/ - Proxy config: compression_interception_params in litellm_settings - Support for input_type param in compress() (openai vs anthropic) - Docs: proxy setup instructions with YAML config example - Tests: 139-line unit test suite for the interception handler Co-Authored-By: Claude Opus 4.6 --- .../docs/completion/prompt_compression.md | 37 ++ litellm/compression/compress.py | 189 +++++++--- .../compression_interception/__init__.py | 9 + .../compression_interception/handler.py | 332 ++++++++++++++++++ litellm/proxy/common_utils/callback_utils.py | 12 + .../integrations/compression_interception.py | 36 ++ scripts/eval_compression.py | 1 + .../test_compression_interception_handler.py | 139 ++++++++ tests/test_litellm/test_compression.py | 131 ++++++- 9 files changed, 823 insertions(+), 63 deletions(-) create mode 100644 litellm/integrations/compression_interception/__init__.py create mode 100644 litellm/integrations/compression_interception/handler.py create mode 100644 litellm/types/integrations/compression_interception.py create mode 100644 tests/test_litellm/integrations/compression_interception/test_compression_interception_handler.py diff --git a/docs/my-website/docs/completion/prompt_compression.md b/docs/my-website/docs/completion/prompt_compression.md index 2d999291af6..4844634f7bf 100644 --- a/docs/my-website/docs/completion/prompt_compression.md +++ b/docs/my-website/docs/completion/prompt_compression.md @@ -19,6 +19,7 @@ messages = [ compressed = litellm.compress( messages=messages, model="gpt-4o", + input_type="openai_chat_completions", compression_trigger=1000, compression_target=500, ) @@ -30,6 +31,41 @@ response = litellm.completion( ) ``` +## Enable On LiteLLM Proxy (`/v1/messages`) + +If you want prompt compression enabled globally for Anthropic Messages traffic (for example, Claude Code via `/v1/messages`), enable the `compression_interception` callback in your proxy config. + +```yaml +litellm_settings: + callbacks: ["compression_interception"] + compression_interception_params: + enabled_providers: ["bedrock", "anthropic"] + compression_trigger: 12000 + compression_target: 8000 +``` + +With this enabled, the proxy: + +- compresses incoming `/v1/messages` payloads above your trigger +- injects `litellm_content_retrieve` when stubs are created +- runs the retrieval tool loop server-side and returns the final response + +Example request: + +```bash +curl -sS "http://localhost:4000/v1/messages" \ + -H "Authorization: Bearer $LITELLM_PROXY_KEY" \ + -H "Content-Type: application/json" \ + -d '{ + "model": "bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", + "max_tokens": 512, + "messages": [ + {"role":"user","content":[{"type":"text","text":"Large context ..."}]}, + {"role":"user","content":[{"type":"text","text":"Question about that context"}]} + ] + }' +``` + ## What It Returns `compress()` returns a dictionary with: @@ -45,6 +81,7 @@ response = litellm.completion( - `messages` (`List[dict]`, required): input conversation messages - `model` (`str`, required): model name used for token counting +- `input_type` (`str`, required): one of `openai_chat_completions` or `anthropic_messages` - `compression_trigger` (`int`, default `200000`): compress only if input token count exceeds this - `compression_target` (`Optional[int]`, default `70% of compression_trigger`): desired post-compression token budget - `embedding_model` (`Optional[str]`): if set, combines BM25 + embedding relevance scoring diff --git a/litellm/compression/compress.py b/litellm/compression/compress.py index 718bc1c45c3..2f3d8eb5b39 100644 --- a/litellm/compression/compress.py +++ b/litellm/compression/compress.py @@ -3,7 +3,7 @@ Main compress() function — orchestrates BM25/embedding scoring, message stubbi and retrieval tool injection. """ -from typing import Any, Dict, List, Optional, Set +from typing import Any, Dict, List, Literal, Optional, Set, Union, cast from litellm.caching.dual_cache import DualCache from litellm.compression.message_stubbing import ( @@ -15,27 +15,100 @@ 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 from litellm.types.compression import CompressedResult +from litellm.types.llms.anthropic import AllAnthropicMessageValues +from litellm.types.llms.openai import AllMessageValues + +CompressionInputType = Literal["openai_chat_completions", "anthropic_messages"] -def _extract_last_user_message(messages: List[dict]) -> str: - """Return the text content of the last user message.""" - for msg in reversed(messages): - if msg.get("role") == "user": - content = msg.get("content", "") - if isinstance(content, str): - return content - if isinstance(content, list): - parts = [] - for part in content: - if isinstance(part, dict) and part.get("type") == "text": - parts.append(part.get("text", "")) - elif isinstance(part, str): - parts.append(part) - return " ".join(parts) +def _extract_text_content(content: Any) -> str: + if isinstance(content, str): + return content + if isinstance(content, list): + parts: List[str] = [] + for part in content: + if isinstance(part, str): + parts.append(part) + continue + if not isinstance(part, dict): + continue + text = part.get("text") + if isinstance(text, str): + parts.append(text) + continue + tool_content = part.get("content") + if isinstance(tool_content, str): + parts.append(tool_content) + elif isinstance(tool_content, list): + for tool_part in tool_content: + if isinstance(tool_part, str): + parts.append(tool_part) + elif isinstance(tool_part, dict): + nested_text = tool_part.get("text") + if isinstance(nested_text, str): + parts.append(nested_text) + return " ".join(parts) return "" -def _get_protected_indices(messages: List[dict]) -> List[int]: +def _normalize_messages_for_compression( + messages: Union[List[AllMessageValues], List[AllAnthropicMessageValues]], +) -> List[Dict[str, Any]]: + normalized_messages: List[Dict[str, Any]] = [] + for message in messages: + msg_dict = cast(Dict[str, Any], message) + normalized_messages.append( + { + "role": msg_dict.get("role", ""), + "content": _extract_text_content(msg_dict.get("content", "")), + } + ) + return normalized_messages + + +def _remap_compressed_messages( + original_messages: Union[List[AllMessageValues], List[AllAnthropicMessageValues]], + normalized_original_messages: List[Dict[str, Any]], + normalized_compressed_messages: List[Dict[str, Any]], + input_type: CompressionInputType, +) -> List[dict]: + if input_type == "openai_chat_completions": + return normalized_compressed_messages + + remapped_messages: List[dict] = [] + for idx, original_message in enumerate(original_messages): + original_msg = cast(Dict[str, Any], original_message) + remapped_message = {**original_msg} + original_content = original_msg.get("content", "") + original_normalized_content = normalized_original_messages[idx].get( + "content", "" + ) + compressed_content = normalized_compressed_messages[idx].get("content", "") + + if isinstance(original_content, list): + if compressed_content == original_normalized_content: + remapped_message["content"] = original_content + else: + remapped_message["content"] = [ + {"type": "text", "text": compressed_content} + ] + else: + remapped_message["content"] = compressed_content + + remapped_messages.append(remapped_message) + + return remapped_messages + + +def _extract_last_user_message(messages: List[Dict[str, Any]]) -> str: + """Return the text content of the last user message.""" + for msg in reversed(messages): + if msg.get("role") == "user": + return _extract_text_content(msg.get("content", "")) + return "" + + +def _get_protected_indices(messages: List[Dict[str, Any]]) -> List[int]: """ Return indices of messages that must never be compressed: - All system messages @@ -87,8 +160,9 @@ def _combine_scores( def compress( - messages: List[dict], + messages: Union[List[AllMessageValues], List[AllAnthropicMessageValues]], model: str, + input_type: CompressionInputType, compression_trigger: int = 200_000, compression_target: Optional[int] = None, embedding_model: Optional[str] = None, @@ -100,16 +174,18 @@ def compress( Messages below ``compression_trigger`` tokens pass through unchanged. Messages above are scored with BM25 (and optionally embeddings), ranked, - and the lowest-relevance messages are replaced with stubs. Originals are + and the lowest-relevance messages are replaced with stubs. Originals are cached and a retrieval tool is injected so the model can recover dropped content on demand. Parameters: messages: The conversation messages to (potentially) compress. model: The LLM model name — used for token counting. + input_type: Source format for messages. + One of: ``openai_chat_completions`` or ``anthropic_messages``. compression_trigger: Only compress if input exceeds this token count. compression_target: Target token count after compression. - Defaults to ``compression_trigger // 2``. + Defaults to ``compression_trigger * 7 // 10``. embedding_model: If provided, use BM25 + embeddings for scoring. If ``None``, BM25 only. embedding_model_params: Optional kwargs forwarded to @@ -121,15 +197,24 @@ def compress( A ``CompressedResult`` dict containing compressed messages, token counts, a cache of original content, and the retrieval tool definition. """ + if input_type not in ("openai_chat_completions", "anthropic_messages"): + raise ValueError( + "Invalid input_type. Expected 'openai_chat_completions' or 'anthropic_messages'." + ) + if compression_target is None: compression_target = compression_trigger * 7 // 10 - original_tokens = token_counter(model=model, messages=messages) + normalized_messages = _normalize_messages_for_compression(messages=messages) + original_tokens = token_counter( + model=model, + messages=cast(List[Any], normalized_messages), + ) # Pass through if below trigger if original_tokens <= compression_trigger: return CompressedResult( - messages=messages, + messages=cast(List[dict], messages), original_tokens=original_tokens, compressed_tokens=original_tokens, compression_ratio=0.0, @@ -138,10 +223,10 @@ def compress( ) # Extract query for relevance scoring - query = _extract_last_user_message(messages) + query = _extract_last_user_message(normalized_messages) # Score each message - bm25_scores = bm25_score_messages(query, messages) + bm25_scores = bm25_score_messages(query, normalized_messages) if embedding_model: from litellm.compression.scoring.embedding_scorer import ( @@ -150,7 +235,7 @@ def compress( emb_scores = embedding_score_messages( query, - messages, + normalized_messages, model=embedding_model, cache=compression_cache, embedding_model_params=embedding_model_params, @@ -161,36 +246,29 @@ def compress( # Sort message indices by score descending ranked_indices = sorted( - range(len(messages)), + range(len(normalized_messages)), key=lambda i: combined_scores[i], reverse=True, ) # Protected messages are never compressed - protected_indices = _get_protected_indices(messages) + 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=messages[i].get("content", "") or "" + model=model, text=normalized_messages[i].get("content", "") or "" ) # 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_content = messages[idx].get("content", "") or "" + msg_content = normalized_messages[idx].get("content", "") or "" msg_tokens = token_counter(model=model, text=msg_content) remaining = compression_target - current_tokens @@ -198,12 +276,10 @@ def compress( 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(messages[idx], remaining) + truncated = truncate_message(normalized_messages[idx], remaining) truncated_tokens = token_counter( model=model, text=truncated.get("content", "") or "", @@ -217,33 +293,36 @@ def compress( cache: Dict[str, str] = {} used_keys: Set[str] = set() - for i, msg in enumerate(messages): + for i, msg in enumerate(normalized_messages): if i in kept_indices: - # 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", "") - if isinstance(content, list): - content = " ".join( - p.get("text", "") if isinstance(p, dict) else str(p) - for p in content - ) - cache[key] = content + cache[key] = msg.get("content", "") compressed_messages.append(stub_message(msg, key)) - # Build retrieval tool - tools = [build_retrieval_tool(list(cache.keys()))] if cache else [] + remapped_compressed_messages = _remap_compressed_messages( + original_messages=messages, + normalized_original_messages=normalized_messages, + normalized_compressed_messages=compressed_messages, + input_type=input_type, + ) - compressed_tokens = token_counter(model=model, messages=compressed_messages) + tools = [build_retrieval_tool(list(cache.keys()))] if cache else [] + compressed_tokens = token_counter( + model=model, + messages=cast(List[Any], compressed_messages), + ) return CompressedResult( - messages=compressed_messages, + messages=remapped_compressed_messages, original_tokens=original_tokens, compressed_tokens=compressed_tokens, - compression_ratio=round(1 - (compressed_tokens / original_tokens), 4) - if original_tokens > 0 - else 0.0, + compression_ratio=( + round(1 - (compressed_tokens / original_tokens), 4) + if original_tokens > 0 + else 0.0 + ), cache=cache, tools=tools, ) diff --git a/litellm/integrations/compression_interception/__init__.py b/litellm/integrations/compression_interception/__init__.py new file mode 100644 index 00000000000..bbcf757cbba --- /dev/null +++ b/litellm/integrations/compression_interception/__init__.py @@ -0,0 +1,9 @@ +""" +Compression Interception integration package. +""" + +from litellm.integrations.compression_interception.handler import ( + CompressionInterceptionLogger, +) + +__all__ = ["CompressionInterceptionLogger"] diff --git a/litellm/integrations/compression_interception/handler.py b/litellm/integrations/compression_interception/handler.py new file mode 100644 index 00000000000..7e33f9e4abb --- /dev/null +++ b/litellm/integrations/compression_interception/handler.py @@ -0,0 +1,332 @@ +""" +Compression Interception Handler for /v1/messages. + +CustomLogger that applies prompt compression in async_pre_request_hook and +executes litellm_content_retrieve tool calls in the Anthropic agentic loop. +""" + +import json +from typing import Any, Dict, List, Optional, Tuple, Union, cast + +import litellm +from litellm._logging import verbose_logger +from litellm.anthropic_interface import messages as anthropic_messages +from litellm.integrations.custom_logger import CustomLogger +from litellm.types.integrations.compression_interception import ( + CompressionInterceptionConfig, +) +from litellm.types.llms.anthropic import AllAnthropicMessageValues +from litellm.types.utils import LlmProviders + +_RETRIEVAL_TOOL_NAME = "litellm_content_retrieve" +_COMPRESSION_CACHE_KEY = "_compression_interception_cache" + + +def _is_retrieval_tool(tool: Dict[str, Any]) -> bool: + if not isinstance(tool, dict): + return False + if tool.get("name") == _RETRIEVAL_TOOL_NAME: + return True + if tool.get("type") == "function": + function_obj = tool.get("function", {}) + if isinstance(function_obj, dict): + return function_obj.get("name") == _RETRIEVAL_TOOL_NAME + return False + + +def _merge_tools( + existing_tools: Optional[List[Dict[str, Any]]], + new_tools: Optional[List[Dict[str, Any]]], +) -> List[Dict[str, Any]]: + merged: List[Dict[str, Any]] = [] + seen_names = set() + + for tool in (existing_tools or []) + (new_tools or []): + if not isinstance(tool, dict): + continue + name = tool.get("name") + if not name and tool.get("type") == "function": + function_obj = tool.get("function", {}) + if isinstance(function_obj, dict): + name = function_obj.get("name") + dedupe_key = name or str(tool) + if dedupe_key in seen_names: + continue + seen_names.add(dedupe_key) + merged.append(tool) + + return merged + + +def _get_cache_from_kwargs(kwargs: Dict[str, Any]) -> Dict[str, str]: + cache = kwargs.get(_COMPRESSION_CACHE_KEY) + if isinstance(cache, dict): + return cast(Dict[str, str], cache) + + litellm_params = kwargs.get("litellm_params") + if isinstance(litellm_params, dict): + nested_cache = litellm_params.get(_COMPRESSION_CACHE_KEY) + if isinstance(nested_cache, dict): + return cast(Dict[str, str], nested_cache) + return {} + + +def _extract_tool_call_key(tool_call: Dict[str, Any]) -> str: + if "input" in tool_call and isinstance(tool_call["input"], dict): + return str(tool_call["input"].get("key", "")) + + function_obj = tool_call.get("function", {}) + if isinstance(function_obj, dict): + arguments = function_obj.get("arguments", {}) + if isinstance(arguments, dict): + return str(arguments.get("key", "")) + if isinstance(arguments, str): + try: + parsed = json.loads(arguments) + if isinstance(parsed, dict): + return str(parsed.get("key", "")) + except json.JSONDecodeError: + return "" + return "" + + +class CompressionInterceptionLogger(CustomLogger): + def __init__( + self, + enabled_providers: Optional[List[Union[LlmProviders, str]]] = None, + compression_trigger: int = 200_000, + compression_target: Optional[int] = None, + embedding_model: Optional[str] = None, + embedding_model_params: Optional[Dict[str, Any]] = None, + ): + super().__init__() + self.enabled_providers = ( + [p.value if isinstance(p, LlmProviders) else p for p in enabled_providers] + if enabled_providers + else None + ) + self.compression_trigger = compression_trigger + self.compression_target = compression_target + self.embedding_model = embedding_model + self.embedding_model_params = embedding_model_params + + def _resolve_provider(self, kwargs: Dict[str, Any], model: str) -> str: + provider = kwargs.get("custom_llm_provider", "") or kwargs.get( + "litellm_params", {} + ).get("custom_llm_provider", "") + if provider: + return str(provider) + try: + _, provider, _, _ = litellm.get_llm_provider(model=model) + return str(provider) + except Exception: + return "" + + def _provider_enabled(self, provider: str) -> bool: + if self.enabled_providers is None: + return True + return provider in self.enabled_providers + + def _apply_compression( + self, model: str, messages: List[Dict[str, Any]], kwargs: Dict[str, Any] + ) -> Optional[Dict[str, Any]]: + result = litellm.compress( + messages=cast(List[AllAnthropicMessageValues], messages), + model=model, + input_type="anthropic_messages", + compression_trigger=self.compression_trigger, + compression_target=self.compression_target, + embedding_model=self.embedding_model, + embedding_model_params=self.embedding_model_params, + ) + + if result["compression_ratio"] <= 0 and not result["cache"]: + return None + + modified_kwargs = dict(kwargs) + modified_kwargs["messages"] = result["messages"] + modified_kwargs["tools"] = _merge_tools( + cast(Optional[List[Dict[str, Any]]], kwargs.get("tools")), + cast(Optional[List[Dict[str, Any]]], result.get("tools")), + ) + modified_kwargs[_COMPRESSION_CACHE_KEY] = result["cache"] + + litellm_params = modified_kwargs.get("litellm_params") + if isinstance(litellm_params, dict): + litellm_params = dict(litellm_params) + else: + litellm_params = {} + litellm_params[_COMPRESSION_CACHE_KEY] = result["cache"] + modified_kwargs["litellm_params"] = litellm_params + + # Reuse existing fake-stream conversion path used by agentic callbacks. + if modified_kwargs.get("stream"): + modified_kwargs["stream"] = False + modified_kwargs["_websearch_interception_converted_stream"] = True + + verbose_logger.debug( + "CompressionInterception: compressed request " + "original_tokens=%s compressed_tokens=%s ratio=%s", + result["original_tokens"], + result["compressed_tokens"], + result["compression_ratio"], + ) + return modified_kwargs + + async def async_pre_request_hook( + self, model: str, messages: List[Dict], kwargs: Dict + ) -> Optional[Dict]: + provider = self._resolve_provider(kwargs, model=model) + if not self._provider_enabled(provider): + return None + + return self._apply_compression( + model=model, + messages=cast(List[Dict[str, Any]], messages), + kwargs=kwargs, + ) + + def _extract_anthropic_tool_calls(self, response: Any) -> List[Dict[str, Any]]: + if isinstance(response, dict): + content = response.get("content", []) or [] + else: + content = getattr(response, "content", []) or [] + + tool_calls: List[Dict[str, Any]] = [] + for block in content: + if not isinstance(block, dict): + continue + if block.get("type") != "tool_use": + continue + if block.get("name") != _RETRIEVAL_TOOL_NAME: + continue + tool_calls.append( + { + "id": block.get("id", ""), + "name": block.get("name", ""), + "input": block.get("input", {}), + } + ) + return tool_calls + + async def async_should_run_agentic_loop( + self, + response: Any, + model: str, + messages: List[Dict], + tools: Optional[List[Dict]], + stream: bool, + custom_llm_provider: str, + kwargs: Dict, + ) -> Tuple[bool, Dict]: + if not self._provider_enabled(custom_llm_provider): + return False, {} + if not any(_is_retrieval_tool(t) for t in (tools or [])): + return False, {} + cache = _get_cache_from_kwargs(kwargs) + if not cache: + return False, {} + + tool_calls = self._extract_anthropic_tool_calls(response) + if not tool_calls: + return False, {} + return True, {"tool_calls": tool_calls, "cache": cache} + + async def async_run_agentic_loop( + self, + tools: Dict, + model: str, + messages: List[Dict], + response: Any, + anthropic_messages_provider_config: Any, + anthropic_messages_optional_request_params: Dict, + logging_obj: Any, + stream: bool, + kwargs: Dict, + ) -> Any: + tool_calls = cast(List[Dict[str, Any]], tools.get("tool_calls", [])) + cache = cast(Dict[str, str], tools.get("cache", {})) + + assistant_blocks: List[Dict[str, Any]] = [] + tool_result_blocks: List[Dict[str, Any]] = [] + + for tool_call in tool_calls: + key = _extract_tool_call_key(tool_call) + tool_call_id = str(tool_call.get("id", "")) + assistant_blocks.append( + { + "type": "tool_use", + "id": tool_call_id, + "name": _RETRIEVAL_TOOL_NAME, + "input": {"key": key}, + } + ) + tool_result_blocks.append( + { + "type": "tool_result", + "tool_use_id": tool_call_id, + "content": cache.get(key, f"[key {key!r} not found in cache]"), + } + ) + + follow_up_messages = messages + [ + {"role": "assistant", "content": assistant_blocks}, + {"role": "user", "content": tool_result_blocks}, + ] + + kwargs_for_followup = { + k: v + for k, v in kwargs.items() + if not k.startswith("_compression_interception") + and not k.startswith("_websearch_interception") + and k != "litellm_logging_obj" + } + optional_params_without_max_tokens = { + k: v + for k, v in anthropic_messages_optional_request_params.items() + if k != "max_tokens" + } + + max_tokens = anthropic_messages_optional_request_params.get( + "max_tokens", kwargs.get("max_tokens", 1024) + ) + full_model_name = model + if logging_obj is not None: + agentic_params = logging_obj.model_call_details.get( + "agentic_loop_params", {} + ) + full_model_name = agentic_params.get("model", model) + + return await anthropic_messages.acreate( + max_tokens=max_tokens, + messages=follow_up_messages, + model=full_model_name, + **optional_params_without_max_tokens, + **kwargs_for_followup, + ) + + @classmethod + def from_config_yaml( + cls, config: CompressionInterceptionConfig + ) -> "CompressionInterceptionLogger": + enabled_providers = config.get("enabled_providers") + return cls( + enabled_providers=enabled_providers, + compression_trigger=config.get("compression_trigger", 200_000), + compression_target=config.get("compression_target"), + embedding_model=config.get("embedding_model"), + embedding_model_params=config.get("embedding_model_params"), + ) + + @staticmethod + def initialize_from_proxy_config( + litellm_settings: Dict[str, Any], + callback_specific_params: Dict[str, Any], + ) -> "CompressionInterceptionLogger": + compression_params: CompressionInterceptionConfig = {} + if "compression_interception_params" in litellm_settings: + compression_params = litellm_settings["compression_interception_params"] + elif "compression_interception" in callback_specific_params: + compression_params = callback_specific_params["compression_interception"] + + return CompressionInterceptionLogger.from_config_yaml(compression_params) diff --git a/litellm/proxy/common_utils/callback_utils.py b/litellm/proxy/common_utils/callback_utils.py index 9ecae363ed7..a4cc503951f 100644 --- a/litellm/proxy/common_utils/callback_utils.py +++ b/litellm/proxy/common_utils/callback_utils.py @@ -281,6 +281,18 @@ def initialize_callbacks_on_proxy( # noqa: PLR0915 ) ) imported_list.append(websearch_interception_obj) + elif isinstance(callback, str) and callback == "compression_interception": + from litellm.integrations.compression_interception.handler import ( + CompressionInterceptionLogger, + ) + + compression_interception_obj = ( + CompressionInterceptionLogger.initialize_from_proxy_config( + litellm_settings=litellm_settings, + callback_specific_params=callback_specific_params, + ) + ) + imported_list.append(compression_interception_obj) elif isinstance(callback, str) and callback == "datadog_cost_management": from litellm.integrations.datadog.datadog_cost_management import ( DatadogCostManagementLogger, diff --git a/litellm/types/integrations/compression_interception.py b/litellm/types/integrations/compression_interception.py new file mode 100644 index 00000000000..49ad062bcb2 --- /dev/null +++ b/litellm/types/integrations/compression_interception.py @@ -0,0 +1,36 @@ +""" +Type definitions for Compression Interception integration. +""" + +from typing import Any, Dict, List, Optional, TypedDict + + +class CompressionInterceptionConfig(TypedDict, total=False): + """ + Configuration parameters for CompressionInterceptionLogger. + + Used in proxy_config.yaml under litellm_settings: + litellm_settings: + compression_interception_params: + enabled_providers: ["openai", "anthropic"] + compression_trigger: 12000 + compression_target: 8000 + embedding_model: "text-embedding-3-small" + embedding_model_params: + dimensions: 512 + """ + + enabled_providers: List[str] + """Optional provider allowlist. If omitted, applies to all providers.""" + + compression_trigger: int + """Only compress requests above this prompt token threshold.""" + + compression_target: Optional[int] + """Target token count after compression. If None, uses compressor default.""" + + embedding_model: Optional[str] + """Optional embedding model used for hybrid ranking.""" + + embedding_model_params: Optional[Dict[str, Any]] + """Optional params forwarded to litellm.embedding() scorer.""" diff --git a/scripts/eval_compression.py b/scripts/eval_compression.py index d7d90dacc2e..c35742f6da6 100644 --- a/scripts/eval_compression.py +++ b/scripts/eval_compression.py @@ -880,6 +880,7 @@ def eval_problem( result = litellm.compress( messages=messages, model=model, + input_type="openai_chat_completions", compression_trigger=compression_trigger, embedding_model=embedding_model, ) diff --git a/tests/test_litellm/integrations/compression_interception/test_compression_interception_handler.py b/tests/test_litellm/integrations/compression_interception/test_compression_interception_handler.py new file mode 100644 index 00000000000..bb7b60f0d79 --- /dev/null +++ b/tests/test_litellm/integrations/compression_interception/test_compression_interception_handler.py @@ -0,0 +1,139 @@ +from unittest.mock import AsyncMock, patch + +import pytest + +from litellm.integrations.compression_interception.handler import ( + CompressionInterceptionLogger, +) + + +def test_initialize_from_proxy_config(): + litellm_settings = { + "compression_interception_params": { + "enabled_providers": ["openai"], + "compression_trigger": 12000, + "compression_target": 8000, + } + } + logger = CompressionInterceptionLogger.initialize_from_proxy_config( + litellm_settings=litellm_settings, callback_specific_params={} + ) + assert logger.enabled_providers == ["openai"] + assert logger.compression_trigger == 12000 + assert logger.compression_target == 8000 + + +@pytest.mark.asyncio +async def test_async_pre_request_hook_applies_compression_and_merges_tools(): + logger = CompressionInterceptionLogger(enabled_providers=["bedrock"]) + kwargs = { + "messages": [{"role": "user", "content": "big context"}], + "tools": [ + {"type": "function", "function": {"name": "calculator", "parameters": {}}} + ], + "custom_llm_provider": "bedrock", + "litellm_params": {"request_id": "abc"}, + "stream": True, + } + mock_result = { + "messages": [{"role": "user", "content": "compressed context"}], + "original_tokens": 2000, + "compressed_tokens": 900, + "compression_ratio": 0.55, + "cache": {"auth.py": "full content"}, + "tools": [ + { + "type": "function", + "function": {"name": "litellm_content_retrieve", "parameters": {}}, + } + ], + } + with patch( + "litellm.integrations.compression_interception.handler.litellm.compress", + return_value=mock_result, + ): + result = await logger.async_pre_request_hook( + model="bedrock/claude-3-5-sonnet", + messages=[{"role": "user", "content": "big context"}], + kwargs=kwargs, + ) + + assert result is not None + assert result["messages"][0]["content"] == "compressed context" + tool_names = [t.get("function", {}).get("name") for t in result["tools"]] + assert "calculator" in tool_names + assert "litellm_content_retrieve" in tool_names + assert result["_compression_interception_cache"]["auth.py"] == "full content" + assert ( + result["litellm_params"]["_compression_interception_cache"]["auth.py"] + == "full content" + ) + assert result["stream"] is False + assert result["_websearch_interception_converted_stream"] is True + + +@pytest.mark.asyncio +async def test_async_should_run_agentic_loop_detects_anthropic_tool_use(): + logger = CompressionInterceptionLogger(enabled_providers=["bedrock"]) + response = { + "content": [ + { + "type": "tool_use", + "id": "toolu_123", + "name": "litellm_content_retrieve", + "input": {"key": "auth.py"}, + } + ] + } + should_run, payload = await logger.async_should_run_agentic_loop( + response=response, + model="bedrock/claude", + messages=[], + tools=[ + { + "type": "function", + "function": {"name": "litellm_content_retrieve", "parameters": {}}, + } + ], + stream=False, + custom_llm_provider="bedrock", + kwargs={"_compression_interception_cache": {"auth.py": "full content"}}, + ) + assert should_run is True + assert payload["tool_calls"][0]["id"] == "toolu_123" + + +@pytest.mark.asyncio +async def test_async_run_agentic_loop_executes_anthropic_followup(): + logger = CompressionInterceptionLogger(enabled_providers=["bedrock"]) + with patch( + "litellm.integrations.compression_interception.handler.anthropic_messages.acreate", + new=AsyncMock(return_value={"final": "answer"}), + ) as mock_acreate: + result = await logger.async_run_agentic_loop( + tools={ + "tool_calls": [ + { + "id": "toolu_123", + "name": "litellm_content_retrieve", + "input": {"key": "auth.py"}, + } + ], + "cache": {"auth.py": "full file content"}, + }, + model="bedrock/claude", + messages=[{"role": "user", "content": "fix auth"}], + response={}, + anthropic_messages_provider_config=None, + anthropic_messages_optional_request_params={"max_tokens": 512}, + logging_obj=None, + stream=False, + kwargs={}, + ) + + assert result == {"final": "answer"} + called_messages = mock_acreate.await_args.kwargs["messages"] + assert called_messages[-1]["role"] == "user" + tool_result = called_messages[-1]["content"][0] + assert tool_result["type"] == "tool_result" + assert tool_result["content"] == "full file content" diff --git a/tests/test_litellm/test_compression.py b/tests/test_litellm/test_compression.py index 13dda0cbcbc..5391baae35f 100644 --- a/tests/test_litellm/test_compression.py +++ b/tests/test_litellm/test_compression.py @@ -149,7 +149,9 @@ def test_retrieval_tool_description_lists_keys(): def test_compress_below_trigger_passthrough(): messages = [{"role": "user", "content": "hello"}] - result = litellm.compress(messages, model="gpt-4o") + result = litellm.compress( + messages, model="gpt-4o", input_type="openai_chat_completions" + ) assert result["messages"] == messages assert result["cache"] == {} assert result["tools"] == [] @@ -178,6 +180,7 @@ def test_compress_above_trigger(): result = litellm.compress( big_messages, model="gpt-4o", + input_type="openai_chat_completions", compression_trigger=1000, compression_target=500, ) @@ -195,7 +198,12 @@ def test_compress_preserves_system_message(): {"role": "user", "content": "Large file content. " * 5000}, {"role": "user", "content": "Fix the bug"}, ] - result = litellm.compress(messages, model="gpt-4o", compression_trigger=1000) + result = litellm.compress( + messages, + model="gpt-4o", + input_type="openai_chat_completions", + compression_trigger=1000, + ) assert result["messages"][0]["role"] == "system" assert "System prompt" in result["messages"][0]["content"] @@ -205,7 +213,12 @@ def test_compress_preserves_last_user_message(): {"role": "user", "content": "Big context " * 5000}, {"role": "user", "content": "Fix the bug in auth.py"}, ] - result = litellm.compress(messages, model="gpt-4o", compression_trigger=1000) + result = litellm.compress( + messages, + model="gpt-4o", + input_type="openai_chat_completions", + compression_trigger=1000, + ) last_user = [m for m in result["messages"] if m["role"] == "user"][-1] assert "Fix the bug in auth.py" in last_user["content"] @@ -216,7 +229,12 @@ def test_compress_preserves_last_assistant_message(): {"role": "assistant", "content": "I'll help with that. " * 2000}, {"role": "user", "content": "Now fix the bug"}, ] - result = litellm.compress(messages, model="gpt-4o", compression_trigger=1000) + result = litellm.compress( + messages, + model="gpt-4o", + input_type="openai_chat_completions", + compression_trigger=1000, + ) assistant_msgs = [m for m in result["messages"] if m["role"] == "assistant"] assert len(assistant_msgs) >= 1 # The last assistant message should be preserved (not stubbed) @@ -229,7 +247,12 @@ def test_cache_keys_match_stubs(): {"role": "user", "content": "# auth.py\n" + "code " * 5000}, {"role": "user", "content": "Fix it"}, ] - result = litellm.compress(messages, model="gpt-4o", compression_trigger=1000) + result = litellm.compress( + messages, + model="gpt-4o", + input_type="openai_chat_completions", + compression_trigger=1000, + ) if result["tools"]: tool_desc = result["tools"][0]["function"]["description"] for key in result["cache"]: @@ -237,16 +260,101 @@ def test_cache_keys_match_stubs(): def test_compress_default_target(): - """compression_target defaults to compression_trigger // 2.""" + """compression_target defaults to compression_trigger * 7 // 10.""" messages = [ {"role": "user", "content": "content " * 5000}, {"role": "user", "content": "query"}, ] - result = litellm.compress(messages, model="gpt-4o", compression_trigger=2000) + result = litellm.compress( + messages, + model="gpt-4o", + input_type="openai_chat_completions", + compression_trigger=2000, + ) # Should have compressed — target = 1000 assert result["compressed_tokens"] <= result["original_tokens"] +def test_compress_anthropic_passthrough_below_trigger(): + messages = [ + {"role": "user", "content": [{"type": "text", "text": "hello from anthropic"}]} + ] + result = litellm.compress( + messages=messages, + model="gpt-4o", + input_type="anthropic_messages", + ) + assert result["messages"] == messages + assert result["cache"] == {} + assert result["tools"] == [] + + +def test_compress_anthropic_returns_anthropic_shape(): + messages = [ + {"role": "user", "content": [{"type": "text", "text": "Large file " * 5000}]}, + {"role": "user", "content": [{"type": "text", "text": "Fix auth"}]}, + ] + result = litellm.compress( + messages=messages, + model="gpt-4o", + input_type="anthropic_messages", + compression_trigger=1000, + ) + assert isinstance(result["messages"][0]["content"], list) + if result["cache"]: + first_block = result["messages"][0]["content"][0] + assert isinstance(first_block, dict) + assert first_block.get("type") == "text" + + +def test_compress_anthropic_preserves_last_user_message(): + messages = [ + {"role": "user", "content": [{"type": "text", "text": "Big context " * 5000}]}, + { + "role": "user", + "content": [{"type": "text", "text": "Fix the auth bug in auth.py"}], + }, + ] + result = litellm.compress( + messages=messages, + model="gpt-4o", + input_type="anthropic_messages", + compression_trigger=1000, + ) + + last_user = [m for m in result["messages"] if m["role"] == "user"][-1] + assert isinstance(last_user["content"], list) + assert "Fix the auth bug in auth.py" in last_user["content"][0]["text"] + + +def test_compress_anthropic_cache_keys_match_retrieval_tool(): + messages = [ + { + "role": "user", + "content": [{"type": "text", "text": "# auth.py\n" + "code " * 5000}], + }, + {"role": "user", "content": [{"type": "text", "text": "Fix it"}]}, + ] + result = litellm.compress( + messages=messages, + model="gpt-4o", + input_type="anthropic_messages", + compression_trigger=1000, + ) + + if result["tools"]: + tool_desc = result["tools"][0]["function"]["description"] + for key in result["cache"]: + assert key in tool_desc + + +def test_compress_requires_input_type(): + with pytest.raises(TypeError): + litellm.compress( + messages=[{"role": "user", "content": "hello"}], model="gpt-4o" + ) + + def test_compress_forwards_embedding_model_params(monkeypatch): captured = {} @@ -269,6 +377,7 @@ def test_compress_forwards_embedding_model_params(monkeypatch): {"role": "user", "content": "Fix auth"}, ], model="gpt-4o", + input_type="openai_chat_completions", compression_trigger=1000, embedding_model="text-embedding-3-small", embedding_model_params={"api_base": "https://example-embeddings.test"}, @@ -326,6 +435,7 @@ def test_embedding_scorer(): {"role": "user", "content": "Fix auth"}, ], model="gpt-4o", + input_type="openai_chat_completions", compression_trigger=1000, embedding_model="text-embedding-3-small", ) @@ -346,7 +456,12 @@ def test_simple_compression(final_user_message, expected_content): {"role": "user", "content": "Unrelated cooking recipes " * 2000}, {"role": "user", "content": final_user_message}, ] - result = litellm.compress(messages, model="gpt-4o", compression_trigger=1000) + result = litellm.compress( + messages, + model="gpt-4o", + input_type="openai_chat_completions", + compression_trigger=1000, + ) print(result["messages"]) if expected_content == "Unrelated cooking recipes ": assert "Unrelated cooking recipes " in result["messages"][1]["content"]