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 <noreply@anthropic.com>
This commit is contained in:
Krrish Dholakia 2026-04-13 11:11:08 -07:00
parent 6d96ce81a8
commit 72bd5cb152
9 changed files with 823 additions and 63 deletions

View file

@ -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

View file

@ -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,
)

View file

@ -0,0 +1,9 @@
"""
Compression Interception integration package.
"""
from litellm.integrations.compression_interception.handler import (
CompressionInterceptionLogger,
)
__all__ = ["CompressionInterceptionLogger"]

View file

@ -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)

View file

@ -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,

View file

@ -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."""

View file

@ -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,
)

View file

@ -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"

View file

@ -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"]