fix: compress - make it work on anthropic input as well

This commit is contained in:
Krrish Dholakia 2026-04-14 11:21:49 -07:00
parent 9c20b8c743
commit d1b9036dbf
15 changed files with 1050 additions and 93 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,
)
@ -45,6 +46,7 @@ response = litellm.completion(
- `messages` (`List[dict]`, required): input conversation messages
- `model` (`str`, required): model name used for token counting
- `input_type` (`Literal["anthropic_messages", "openai_chat_completions"]`, required): input message schema
- `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
@ -70,6 +72,28 @@ args = json.loads(tool_call.function.arguments)
full_content = compressed["cache"][args["key"]]
```
## Server-side Callback Loop (`/v1/messages`)
You can enable callback-based compression interception to make retrieval loops
transparent for Anthropic Messages calls:
```yaml
litellm_settings:
callbacks: ["compression_interception"]
compression_interception_params:
enabled: true
compression_trigger: 10000
compression_target: 7000
```
With this enabled, LiteLLM runs the following server-side flow:
1. Compresses inbound messages before the first provider call.
2. Injects the `litellm_content_retrieve` tool.
3. Detects retrieval `tool_use` blocks in the model response.
4. Resolves retrieval keys from the compression cache.
5. Reruns the model via agentic loop and returns the final answer.
## Performance
Benchmarked on [SWE-bench Lite](https://huggingface.co/datasets/princeton-nlp/SWE-bench_Lite_bm25_27K) (real GitHub issues with ~27k tokens of BM25-retrieved repo context per problem).

View file

@ -148,6 +148,7 @@ _custom_logger_compatible_callbacks_literal = Literal[
"vantage",
"posthog",
"levo",
"compression_interception",
]
cold_storage_custom_logger: Optional[_custom_logger_compatible_callbacks_literal] = None
logged_real_time_event_types: Optional[Union[List[str], Literal["*"]]] = None

View file

@ -1,9 +1,9 @@
"""
Main compress() function — orchestrates BM25/embedding scoring, message stubbing,
and retrieval tool injection.
Main compress() function — normalizes input messages, orchestrates BM25/embedding
scoring, message stubbing, and retrieval tool injection.
"""
from typing import Any, Dict, List, Optional, Set
from typing import Any, Dict, List, Optional, Set, Tuple, cast
from litellm.caching.dual_cache import DualCache
from litellm.compression.message_stubbing import (
@ -14,24 +14,94 @@ from litellm.compression.message_stubbing import (
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.compression import CompressedResult, CompressionInputType
def _build_retrieval_tools(keys: List[str], input_type: CompressionInputType) -> List[dict]:
"""
Build retrieval tool definitions in the target request schema.
- OpenAI chat completions: keep OpenAI function-tool schema.
- Anthropic messages: remap OpenAI function-tool schema to Anthropic custom tool.
"""
if not keys:
return []
openai_tools = [build_retrieval_tool(keys)]
if input_type == "openai_chat_completions":
return openai_tools
if input_type == "anthropic_messages":
# Lazy import to avoid introducing provider transformation imports
# during module import for non-Anthropic call paths.
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
anthropic_tools, _mcp_servers = AnthropicConfig()._map_tools(openai_tools)
return cast(List[dict], anthropic_tools)
return openai_tools
def _content_to_text(content: Any) -> str:
"""
Convert OpenAI/Anthropic message content blocks to plain text.
Text extraction policy:
- Include text-bearing fields only (`text` blocks + string values).
- For `tool_result`, recurse into nested `content`.
- Ignore non-textual blocks (images/documents/tool metadata/thinking metadata).
"""
if isinstance(content, str):
return content
if isinstance(content, list):
parts: List[str] = []
for part in content:
if isinstance(part, dict):
part_type = part.get("type")
if part_type == "text":
parts.append(str(part.get("text", "")))
elif part_type == "tool_result":
parts.append(_content_to_text(part.get("content", "")))
elif isinstance(part, str):
parts.append(part)
return " ".join(parts)
return ""
def _normalize_messages_for_compression(
messages: List[dict],
input_type: CompressionInputType,
) -> Tuple[List[dict], List[dict]]:
"""
Normalize each original message to a text-surrogate content for scoring.
Returns:
(normalized_messages, original_messages_copy)
"""
if input_type not in ("anthropic_messages", "openai_chat_completions"):
raise ValueError(
f"Unsupported input_type={input_type}. "
"Expected 'anthropic_messages' or 'openai_chat_completions'."
)
original_messages: List[Dict[str, Any]] = [dict(m) for m in messages]
normalized_messages: List[dict] = []
for msg in original_messages:
normalized_messages.append(
{
**msg,
"content": _content_to_text(msg.get("content", "")),
}
)
return normalized_messages, original_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)
return _content_to_text(msg.get("content", ""))
return ""
@ -89,6 +159,7 @@ def _combine_scores(
def compress(
messages: List[dict],
model: str,
input_type: CompressionInputType = "openai_chat_completions",
compression_trigger: int = 200_000,
compression_target: Optional[int] = None,
embedding_model: Optional[str] = None,
@ -107,6 +178,10 @@ def compress(
Parameters:
messages: The conversation messages to (potentially) compress.
model: The LLM model name — used for token counting.
input_type: Message format of input messages. Must be either:
- "anthropic_messages"
- "openai_chat_completions"
Defaults to "openai_chat_completions" for backward compatibility.
compression_trigger: Only compress if input exceeds this token count.
compression_target: Target token count after compression.
Defaults to ``compression_trigger // 2``.
@ -121,15 +196,23 @@ def compress(
A ``CompressedResult`` dict containing compressed messages, token
counts, a cache of original content, and the retrieval tool definition.
"""
normalized_messages, original_messages = _normalize_messages_for_compression(
messages=messages,
input_type=input_type,
)
if compression_target is None:
compression_target = compression_trigger * 7 // 10
original_tokens = token_counter(model=model, messages=messages)
original_tokens = token_counter(
model=model,
messages=cast(List[Any], original_messages),
)
# Pass through if below trigger
if original_tokens <= compression_trigger:
return CompressedResult(
messages=messages,
messages=original_messages,
original_tokens=original_tokens,
compressed_tokens=original_tokens,
compression_ratio=0.0,
@ -138,10 +221,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 +233,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,20 +244,21 @@ 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=cast(str, normalized_messages[i].get("content", "") or ""),
)
# Fill token budget from highest-scoring messages.
@ -190,8 +274,10 @@ def compress(
for idx in ranked_indices:
if idx in kept_indices:
continue
msg_content = messages[idx].get("content", "") or ""
msg_tokens = token_counter(model=model, text=msg_content)
msg_tokens = token_counter(
model=model,
text=cast(str, normalized_messages[idx].get("content", "") or ""),
)
remaining = compression_target - current_tokens
if remaining <= 0:
@ -203,7 +289,7 @@ def compress(
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(original_messages[idx], remaining)
truncated_tokens = token_counter(
model=model,
text=truncated.get("content", "") or "",
@ -217,33 +303,35 @@ def compress(
cache: Dict[str, str] = {}
used_keys: Set[str] = set()
for i, msg in enumerate(messages):
for i, msg in enumerate(original_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
)
key = extract_key(
normalized_messages[i], fallback_index=i, used_keys=used_keys
)
content = _content_to_text(msg.get("content", ""))
cache[key] = content
compressed_messages.append(stub_message(msg, key))
# Build retrieval tool
tools = [build_retrieval_tool(list(cache.keys()))] if cache else []
# Build retrieval tool in the target request schema
tools = _build_retrieval_tools(list(cache.keys()), input_type=input_type)
compressed_tokens = token_counter(model=model, messages=compressed_messages)
compressed_tokens = token_counter(
model=model,
messages=cast(List[Any], compressed_messages),
)
return CompressedResult(
messages=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,14 @@
"""
Compression Interception Module
Provides server-side prompt compression + retrieval tool fulfillment for
Anthropic Messages agentic loops.
"""
from litellm.integrations.compression_interception.handler import (
CompressionInterceptionLogger,
)
__all__ = [
"CompressionInterceptionLogger",
]

View file

@ -0,0 +1,382 @@
"""
Compression Interception Handler
CustomLogger that compresses inbound Anthropic Messages requests and fulfills
litellm_content_retrieve tool calls server-side via the typed agentic loop plan.
"""
import time
import uuid
from typing import Any, Dict, List, Optional, Tuple, cast
import litellm
from litellm._logging import verbose_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.integrations.compression_interception import (
CompressionInterceptionConfig,
)
from litellm.types.integrations.custom_logger import (
AgenticLoopPlan,
AgenticLoopRequestPatch,
)
from litellm.types.utils import CallTypes
LITELLM_CONTENT_RETRIEVE_TOOL_NAME = "litellm_content_retrieve"
_CACHE_TTL_SECONDS = 15 * 60
class CompressionInterceptionLogger(CustomLogger):
"""
CustomLogger that implements transparent prompt compression + retrieval loops.
Flow:
1. Compress inbound /v1/messages requests in pre-call hook.
2. Inject litellm_content_retrieve tool and persist compressed cache by call_id.
3. Detect retrieval tool_use blocks in first model response.
4. Build typed rerun plan with tool_result blocks from the compressed cache.
"""
def __init__(
self,
enabled: bool = True,
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 = enabled
self.compression_trigger = compression_trigger
self.compression_target = compression_target
self.embedding_model = embedding_model
self.embedding_model_params = embedding_model_params
self._compression_cache_by_call_id: Dict[str, Tuple[Dict[str, str], float]] = {}
@classmethod
def from_config_yaml(
cls, config: CompressionInterceptionConfig
) -> "CompressionInterceptionLogger":
return cls(
enabled=bool(config.get("enabled", True)),
compression_trigger=int(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)
async def async_pre_call_deployment_hook(
self, kwargs: Dict[str, Any], call_type: Optional[CallTypes]
) -> Optional[dict]:
if not self.enabled:
return None
if call_type is not None and call_type != CallTypes.anthropic_messages:
return None
if int(kwargs.get("_agentic_loop_depth", 0) or 0) > 0:
return None
messages = kwargs.get("messages")
model = kwargs.get("model")
if not isinstance(messages, list) or not isinstance(model, str):
return None
if self._has_retrieval_tool(kwargs.get("tools")):
return None
self._prune_expired_cache()
compressed = litellm.compress(
messages=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,
)
kwargs["messages"] = compressed["messages"]
kwargs["tools"] = self._merge_tools(
existing_tools=cast(Optional[List[Dict[str, Any]]], kwargs.get("tools")),
compressed_tools=cast(List[Dict[str, Any]], compressed.get("tools", [])),
)
cache = cast(Dict[str, str], compressed.get("cache", {}))
if cache:
call_id = cast(Optional[str], kwargs.get("litellm_call_id"))
if not call_id:
call_id = str(uuid.uuid4())
kwargs["litellm_call_id"] = call_id
self._compression_cache_by_call_id[call_id] = (cache, time.time())
verbose_logger.debug(
"CompressionInterception: compressed request [call_id=%s original=%d compressed=%d cached_keys=%d]",
call_id,
compressed.get("original_tokens"),
compressed.get("compressed_tokens"),
len(cache),
)
return kwargs
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.enabled:
return False, {}
if not self._has_retrieval_tool(tools):
return False, {}
tool_calls, thinking_blocks = self._extract_retrieval_tool_calls(
response=response
)
if not tool_calls:
return False, {}
return True, {
"tool_calls": tool_calls,
"thinking_blocks": thinking_blocks,
"tool_type": "compression_retrieval",
}
async def async_build_agentic_loop_plan(
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,
) -> AgenticLoopPlan:
self._prune_expired_cache()
tool_calls = cast(List[Dict[str, Any]], tools.get("tool_calls", []))
thinking_blocks = cast(List[Dict[str, Any]], tools.get("thinking_blocks", []))
call_id = self._resolve_call_id(logging_obj=logging_obj, kwargs=kwargs)
cache = self._get_cache(call_id=call_id)
retrieval_results = [
self._resolve_retrieval_content(tc, cache) for tc in tool_calls
]
assistant_message = {
"role": "assistant",
"content": thinking_blocks
+ [
{
"type": "tool_use",
"id": tc.get("id"),
"name": tc.get("name", LITELLM_CONTENT_RETRIEVE_TOOL_NAME),
"input": tc.get("input", {}),
}
for tc in tool_calls
],
}
user_message = {
"role": "user",
"content": [
{
"type": "tool_result",
"tool_use_id": tool_calls[i].get("id"),
"content": retrieval_results[i],
}
for i in range(len(tool_calls))
],
}
follow_up_messages = messages + [assistant_message, user_message]
max_tokens = cast(
Optional[int],
anthropic_messages_optional_request_params.get("max_tokens")
or kwargs.get("max_tokens"),
)
optional_params_without_max_tokens = {
k: v
for k, v in anthropic_messages_optional_request_params.items()
if k != "max_tokens"
}
full_model_name = model
if logging_obj is not None:
agentic_params = logging_obj.model_call_details.get(
"agentic_loop_params", {}
)
full_model_name = cast(str, agentic_params.get("model", model))
request_patch = AgenticLoopRequestPatch(
model=full_model_name,
messages=follow_up_messages,
max_tokens=max_tokens,
optional_params=optional_params_without_max_tokens,
kwargs=self._prepare_followup_kwargs(kwargs=kwargs),
)
return AgenticLoopPlan(
run_agentic_loop=True,
request_patch=request_patch,
metadata={"tool_type": "compression_retrieval", "call_id": call_id or ""},
)
def _prune_expired_cache(self) -> None:
now = time.time()
self._compression_cache_by_call_id = {
call_id: (cache, created_at)
for call_id, (
cache,
created_at,
) in self._compression_cache_by_call_id.items()
if now - created_at <= _CACHE_TTL_SECONDS
}
def _get_cache(self, call_id: Optional[str]) -> Dict[str, str]:
if not call_id:
return {}
cache_entry = self._compression_cache_by_call_id.get(call_id)
if cache_entry is None:
return {}
return cache_entry[0]
def _resolve_call_id(
self, logging_obj: Any, kwargs: Dict[str, Any]
) -> Optional[str]:
if logging_obj is not None:
logging_call_id = getattr(logging_obj, "litellm_call_id", None)
if isinstance(logging_call_id, str) and logging_call_id:
return logging_call_id
kwargs_call_id = kwargs.get("litellm_call_id")
return cast(
Optional[str], kwargs_call_id if isinstance(kwargs_call_id, str) else None
)
def _resolve_retrieval_content(
self, tool_call: Dict[str, Any], cache: Dict[str, str]
) -> str:
raw_input = tool_call.get("input", {})
key = ""
if isinstance(raw_input, dict):
key = str(raw_input.get("key", "") or "")
if not key:
return "No retrieval key provided."
if key in cache:
return cache[key]
return f"[compressed content key '{key}' not found]"
def _extract_retrieval_tool_calls(
self, response: Any
) -> Tuple[List[Dict[str, Any]], List[Dict[str, Any]]]:
if isinstance(response, dict):
content = response.get("content", [])
else:
content = getattr(response, "content", []) or []
if not isinstance(content, list):
return [], []
tool_calls: List[Dict[str, Any]] = []
thinking_blocks: List[Dict[str, Any]] = []
for block in content:
if isinstance(block, dict):
block_type = block.get("type")
block_name = block.get("name")
if block_type in ("thinking", "redacted_thinking"):
thinking_blocks.append(block)
if (
block_type == "tool_use"
and block_name == LITELLM_CONTENT_RETRIEVE_TOOL_NAME
):
tool_calls.append(
{
"id": block.get("id"),
"type": "tool_use",
"name": block_name,
"input": block.get("input", {}),
}
)
else:
block_type = getattr(block, "type", None)
block_name = getattr(block, "name", None)
if block_type == "thinking":
thinking_blocks.append(
{
"type": "thinking",
"thinking": getattr(block, "thinking", ""),
"signature": getattr(block, "signature", ""),
}
)
elif block_type == "redacted_thinking":
thinking_blocks.append(
{
"type": "redacted_thinking",
"data": getattr(block, "data", ""),
}
)
if (
block_type == "tool_use"
and block_name == LITELLM_CONTENT_RETRIEVE_TOOL_NAME
):
tool_calls.append(
{
"id": getattr(block, "id", None),
"type": "tool_use",
"name": block_name,
"input": getattr(block, "input", {}) or {},
}
)
return tool_calls, thinking_blocks
def _prepare_followup_kwargs(self, kwargs: Dict[str, Any]) -> Dict[str, Any]:
internal_keys = {"litellm_logging_obj"}
return {
k: v
for k, v in kwargs.items()
if not k.startswith("_compression_interception") and k not in internal_keys
}
def _has_retrieval_tool(self, tools: Any) -> bool:
if not isinstance(tools, list):
return False
for tool in tools:
if not isinstance(tool, dict):
continue
function = tool.get("function")
if tool.get("type") == "function" and isinstance(function, dict):
if function.get("name") == LITELLM_CONTENT_RETRIEVE_TOOL_NAME:
return True
if (
tool.get("type") == "custom"
and tool.get("name") == LITELLM_CONTENT_RETRIEVE_TOOL_NAME
):
return True
return False
def _merge_tools(
self,
existing_tools: Optional[List[Dict[str, Any]]],
compressed_tools: List[Dict[str, Any]],
) -> List[Dict[str, Any]]:
merged = list(existing_tools or [])
if self._has_retrieval_tool(merged):
return merged
merged.extend(compressed_tools)
return merged

View file

@ -4488,7 +4488,9 @@ class BaseLLMHTTPHandler:
full_model_name = model
if logging_obj is not None:
agentic_params = logging_obj.model_call_details.get("agentic_loop_params", {})
agentic_params = logging_obj.model_call_details.get(
"agentic_loop_params", {}
)
full_model_name = cast(str, agentic_params.get("model", model))
optional_params = dict(anthropic_messages_optional_request_params)
@ -4508,7 +4510,9 @@ class BaseLLMHTTPHandler:
kwargs_for_followup = {
k: v
for k, v in kwargs.items()
if not k.startswith("_websearch_interception") and k not in internal_keys
if not k.startswith("_websearch_interception")
and not k.startswith("_compression_interception")
and k not in internal_keys
}
kwargs_for_followup.update(patch.kwargs)
kwargs_for_followup["_agentic_loop_depth"] = depth + 1
@ -4562,6 +4566,7 @@ class BaseLLMHTTPHandler:
k: v
for k, v in kwargs.items()
if not k.startswith("_websearch_interception")
and not k.startswith("_compression_interception")
and k not in internal_params
}
kwargs_for_followup.update(patch.kwargs)
@ -4632,7 +4637,9 @@ class BaseLLMHTTPHandler:
)
kwargs_with_provider = kwargs.copy() if kwargs else {}
kwargs_with_provider["custom_llm_provider"] = custom_llm_provider
kwargs_with_provider[
"custom_llm_provider"
] = custom_llm_provider
build_plan_overridden = (
callback.__class__.async_build_agentic_loop_plan
is not CustomLogger.async_build_agentic_loop_plan
@ -4797,21 +4804,25 @@ class BaseLLMHTTPHandler:
)
kwargs_with_provider = kwargs.copy() if kwargs else {}
kwargs_with_provider["custom_llm_provider"] = custom_llm_provider
kwargs_with_provider[
"custom_llm_provider"
] = custom_llm_provider
build_plan_overridden = (
callback.__class__.async_build_chat_completion_agentic_loop_plan
is not CustomLogger.async_build_chat_completion_agentic_loop_plan
)
if not build_plan_overridden:
return await callback.async_run_chat_completion_agentic_loop(
tools=tool_calls,
model=model,
messages=messages,
response=response,
optional_params=optional_params,
logging_obj=logging_obj,
stream=stream,
kwargs=kwargs_with_provider,
return (
await callback.async_run_chat_completion_agentic_loop(
tools=tool_calls,
model=model,
messages=messages,
response=response,
optional_params=optional_params,
logging_obj=logging_obj,
stream=stream,
kwargs=kwargs_with_provider,
)
)
plan = await callback.async_build_chat_completion_agentic_loop_plan(

View file

@ -22,11 +22,21 @@ model_list:
output_cost_per_token: 10 # 100x standard ($10.00/1M = $0.00001)
# Anthropic model for /v1/messages test — 100x custom pricing
- model_name: "claude-sonnet-4-20250514"
- model_name: "claude-sonnet-4-6"
litellm_params:
model: anthropic/claude-sonnet-4-20250514
model: anthropic/claude-sonnet-4-6
api_key: os.environ/ANTHROPIC_API_KEY
model_info:
id: claude-sonnet-4-custom-pricing
input_cost_per_token: 0.0003 # 100x standard ($0.000003)
output_cost_per_token: 0.0015 # 100x standard ($0.000015)
output_cost_per_token: 0.0015 # 100x standard ($0.000015)
litellm_settings:
callbacks: ["compression_interception"]
compression_interception_params:
enabled: true
compression_trigger: 1000
# optional:
# embedding_model: "text-embedding-3-small"
# embedding_model_params:
# dimensions: 512

View file

@ -37,6 +37,20 @@ def initialize_callbacks_on_proxy( # noqa: PLR0915
if isinstance(value, list):
imported_list: List[Any] = []
for callback in value: # ["presidio", <my-custom-callback>]
if 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)
continue
# check if callback is a custom logger compatible callback
if isinstance(callback, str):
callback = LoggingCallbackManager._add_custom_callback_generic_api_str(

View file

@ -2,7 +2,9 @@
Type definitions for litellm.compress().
"""
from typing import Dict, List, TypedDict
from typing import Dict, List, Literal, TypedDict
CompressionInputType = Literal["anthropic_messages", "openai_chat_completions"]
class CompressedResult(TypedDict):

View file

@ -0,0 +1,27 @@
"""
Type definitions for Compression Interception integration.
"""
from typing import Any, Dict, 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: true
compression_trigger: 100000
compression_target: 70000
embedding_model: "text-embedding-3-small"
embedding_model_params:
dimensions: 512
"""
enabled: bool
compression_trigger: int
compression_target: Optional[int]
embedding_model: Optional[str]
embedding_model_params: Optional[Dict[str, Any]]

View file

@ -781,9 +781,9 @@ def function_setup( # noqa: PLR0915
coroutine_checker = get_coroutine_checker_fn()
## DYNAMIC CALLBACKS ##
dynamic_callbacks: Optional[
List[Union[str, Callable, "CustomLogger"]]
] = kwargs.pop("callbacks", None)
dynamic_callbacks: Optional[List[Union[str, Callable, "CustomLogger"]]] = (
kwargs.pop("callbacks", None)
)
all_callbacks = get_dynamic_callbacks(dynamic_callbacks=dynamic_callbacks)
if len(all_callbacks) > 0:
@ -1689,9 +1689,9 @@ def client(original_function): # noqa: PLR0915
exception=e,
retry_policy=kwargs.get("retry_policy"),
)
kwargs[
"retry_policy"
] = reset_retry_policy() # prevent infinite loops
kwargs["retry_policy"] = (
reset_retry_policy()
) # prevent infinite loops
litellm.num_retries = (
None # set retries to None to prevent infinite loops
)
@ -1738,9 +1738,9 @@ def client(original_function): # noqa: PLR0915
exception=e,
retry_policy=kwargs.get("retry_policy"),
)
kwargs[
"retry_policy"
] = reset_retry_policy() # prevent infinite loops
kwargs["retry_policy"] = (
reset_retry_policy()
) # prevent infinite loops
litellm.num_retries = (
None # set retries to None to prevent infinite loops
)
@ -3774,10 +3774,10 @@ def pre_process_non_default_params(
if "response_format" in non_default_params:
if provider_config is not None:
non_default_params[
"response_format"
] = provider_config.get_json_schema_from_pydantic_object(
response_format=non_default_params["response_format"]
non_default_params["response_format"] = (
provider_config.get_json_schema_from_pydantic_object(
response_format=non_default_params["response_format"]
)
)
else:
non_default_params["response_format"] = type_to_response_format_param(
@ -3906,16 +3906,16 @@ def pre_process_optional_params(
True # so that main.py adds the function call to the prompt
)
if "tools" in non_default_params:
optional_params[
"functions_unsupported_model"
] = non_default_params.pop("tools")
optional_params["functions_unsupported_model"] = (
non_default_params.pop("tools")
)
non_default_params.pop(
"tool_choice", None
) # causes ollama requests to hang
elif "functions" in non_default_params:
optional_params[
"functions_unsupported_model"
] = non_default_params.pop("functions")
optional_params["functions_unsupported_model"] = (
non_default_params.pop("functions")
)
elif (
litellm.add_function_to_prompt
): # if user opts to add it to prompt instead
@ -4896,9 +4896,7 @@ def _get_order_filtered_deployments(
) -> List:
if target_order is not None:
filtered = [
d
for d in healthy_deployments
if _get_deployment_order(d) == target_order
d for d in healthy_deployments if _get_deployment_order(d) == target_order
]
if filtered:
return filtered
@ -5875,8 +5873,12 @@ def _get_model_info_helper( # noqa: PLR0915
supports_web_search=_model_info.get("supports_web_search", None),
supports_url_context=_model_info.get("supports_url_context", None),
supports_reasoning=_model_info.get("supports_reasoning", None),
supports_none_reasoning_effort=_model_info.get("supports_none_reasoning_effort", None),
supports_xhigh_reasoning_effort=_model_info.get("supports_xhigh_reasoning_effort", None),
supports_none_reasoning_effort=_model_info.get(
"supports_none_reasoning_effort", None
),
supports_xhigh_reasoning_effort=_model_info.get(
"supports_xhigh_reasoning_effort", None
),
supports_computer_use=_model_info.get("supports_computer_use", None),
search_context_cost_per_query=_model_info.get(
"search_context_cost_per_query", None
@ -7554,9 +7556,9 @@ class ModelResponseIterator:
if convert_to_delta is True:
_stream_response = ModelResponseStream()
_stream_response.choices[0].delta.content = model_response.choices[0].message.content # type: ignore
self.model_response: Union[
ModelResponse, ModelResponseStream
] = _stream_response
self.model_response: Union[ModelResponse, ModelResponseStream] = (
_stream_response
)
else:
self.model_response = model_response
self.is_done = False

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,217 @@
"""
Unit tests for Compression Interception Handler.
"""
from unittest.mock import MagicMock
import pytest
from litellm.integrations.compression_interception.handler import (
CompressionInterceptionLogger,
)
from litellm.types.utils import CallTypes
def test_initialize_from_proxy_config():
"""Test initialization from proxy config with litellm_settings."""
litellm_settings = {
"compression_interception_params": {
"enabled": True,
"compression_trigger": 1234,
"compression_target": 789,
}
}
logger = CompressionInterceptionLogger.initialize_from_proxy_config(
litellm_settings=litellm_settings,
callback_specific_params={},
)
assert logger.enabled is True
assert logger.compression_trigger == 1234
assert logger.compression_target == 789
@pytest.mark.asyncio
async def test_pre_call_hook_compresses_messages_and_injects_tool(monkeypatch):
"""Test pre-call hook compresses and stores per-call cache."""
logger = CompressionInterceptionLogger()
compressed_result = {
"messages": [{"role": "user", "content": "stubbed"}],
"original_tokens": 12000,
"compressed_tokens": 5000,
"compression_ratio": 0.58,
"cache": {"auth.py": "full file content"},
"tools": [
{
"type": "function",
"function": {
"name": "litellm_content_retrieve",
"parameters": {
"type": "object",
"properties": {"key": {"type": "string"}},
},
},
}
],
}
def _fake_compress(**kwargs):
return compressed_result
monkeypatch.setattr("litellm.compress", _fake_compress)
kwargs = {
"model": "bedrock/us.anthropic.claude-sonnet-4-5",
"messages": [{"role": "user", "content": "very large context"}],
"tools": [
{
"type": "function",
"function": {"name": "existing_tool", "parameters": {"type": "object"}},
}
],
}
result = await logger.async_pre_call_deployment_hook(
kwargs=kwargs, call_type=CallTypes.anthropic_messages
)
assert result is not None
assert result["messages"] == compressed_result["messages"]
tool_names = [t.get("function", {}).get("name") for t in result["tools"]]
assert "existing_tool" in tool_names
assert "litellm_content_retrieve" in tool_names
assert result["litellm_call_id"] in logger._compression_cache_by_call_id
@pytest.mark.asyncio
async def test_should_run_agentic_loop_detects_retrieval_tool_use():
"""Test should-run hook returns tool calls for retrieval tool_use blocks."""
logger = CompressionInterceptionLogger()
response = {
"content": [
{
"type": "tool_use",
"id": "toolu_123",
"name": "litellm_content_retrieve",
"input": {"key": "auth.py"},
}
]
}
should_run, tools_dict = await logger.async_should_run_agentic_loop(
response=response,
model="bedrock/claude",
messages=[],
tools=[
{
"type": "function",
"function": {
"name": "litellm_content_retrieve",
"parameters": {"type": "object"},
},
}
],
stream=False,
custom_llm_provider="bedrock",
kwargs={},
)
assert should_run is True
assert len(tools_dict["tool_calls"]) == 1
assert tools_dict["tool_calls"][0]["input"]["key"] == "auth.py"
@pytest.mark.asyncio
async def test_build_agentic_loop_plan_returns_request_patch():
"""Callback should return typed patch with tool_result content."""
logger = CompressionInterceptionLogger()
call_id = "call_123"
logger._compression_cache_by_call_id[call_id] = (
{"auth.py": "full auth file"},
9999999999.0,
)
logging_obj = MagicMock()
logging_obj.litellm_call_id = call_id
logging_obj.model_call_details = {
"agentic_loop_params": {"model": "bedrock/invoke/claude-3-5-sonnet"}
}
plan = await logger.async_build_agentic_loop_plan(
tools={
"tool_calls": [
{
"id": "toolu_abc",
"type": "tool_use",
"name": "litellm_content_retrieve",
"input": {"key": "auth.py"},
}
]
},
model="claude-3-5-sonnet",
messages=[{"role": "user", "content": "read auth.py"}],
response=None,
anthropic_messages_provider_config=None,
anthropic_messages_optional_request_params={
"max_tokens": 1024,
"tools": [{"name": "litellm_content_retrieve"}],
},
logging_obj=logging_obj,
stream=False,
kwargs={
"temperature": 0.1,
"_compression_interception_internal": True,
"litellm_logging_obj": object(),
},
)
assert plan.run_agentic_loop is True
assert plan.request_patch is not None
assert plan.request_patch.model == "bedrock/invoke/claude-3-5-sonnet"
assert plan.request_patch.max_tokens == 1024
assert plan.request_patch.messages is not None
assert len(plan.request_patch.messages) == 3
tool_result_content = plan.request_patch.messages[-1]["content"][0]["content"]
assert tool_result_content == "full auth file"
assert "_compression_interception_internal" not in plan.request_patch.kwargs
assert "litellm_logging_obj" not in plan.request_patch.kwargs
assert plan.request_patch.kwargs["temperature"] == 0.1
assert "max_tokens" not in plan.request_patch.optional_params
@pytest.mark.asyncio
async def test_build_agentic_loop_plan_missing_key_fallback():
"""Missing cache keys should produce deterministic fallback content."""
logger = CompressionInterceptionLogger()
logging_obj = MagicMock()
logging_obj.litellm_call_id = "missing_call"
logging_obj.model_call_details = {"agentic_loop_params": {}}
plan = await logger.async_build_agentic_loop_plan(
tools={
"tool_calls": [
{
"id": "toolu_missing",
"type": "tool_use",
"name": "litellm_content_retrieve",
"input": {"key": "not_found.py"},
}
]
},
model="claude-3-5-sonnet",
messages=[{"role": "user", "content": "read file"}],
response=None,
anthropic_messages_provider_config=None,
anthropic_messages_optional_request_params={},
logging_obj=logging_obj,
stream=False,
kwargs={},
)
assert plan.request_patch is not None
assert (
plan.request_patch.messages[-1]["content"][0]["content"]
== "[compressed content key 'not_found.py' not found]"
)

View file

@ -1,14 +1,17 @@
import sys
import os
from types import SimpleNamespace
sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
from litellm.proxy.common_utils.callback_utils import (
initialize_callbacks_on_proxy,
get_remaining_tokens_and_requests_from_request_data,
normalize_callback_names,
)
import litellm
from unittest.mock import patch
from litellm.proxy.common_utils.callback_utils import process_callback
@ -79,5 +82,40 @@ def test_normalize_callback_names_none_returns_empty_list():
def test_normalize_callback_names_lowercases_strings():
assert normalize_callback_names(["SQS", "S3", "CUSTOM_CALLBACK"]) == ["sqs", "s3", "custom_callback"]
assert normalize_callback_names(["SQS", "S3", "CUSTOM_CALLBACK"]) == [
"sqs",
"s3",
"custom_callback",
]
def test_initialize_callbacks_on_proxy_instantiates_compression_interception(
monkeypatch,
):
dummy_callback = object()
monkeypatch.setitem(
sys.modules,
"litellm.proxy.proxy_server",
SimpleNamespace(prisma_client=None),
)
monkeypatch.setattr(
"litellm.integrations.compression_interception.handler.CompressionInterceptionLogger.initialize_from_proxy_config",
lambda litellm_settings, callback_specific_params: dummy_callback,
)
original_callbacks = (
list(litellm.callbacks) if isinstance(litellm.callbacks, list) else []
)
litellm.callbacks = []
try:
initialize_callbacks_on_proxy(
value=["compression_interception"],
premium_user=False,
config_file_path=".",
litellm_settings={"compression_interception_params": {"enabled": True}},
callback_specific_params={},
)
assert dummy_callback in litellm.callbacks
assert "compression_interception" not in litellm.callbacks
finally:
litellm.callbacks = original_callbacks

View file

@ -13,6 +13,9 @@ from litellm.compression.content_detection import detect_content_type
from litellm.compression.message_stubbing import extract_key, stub_message
from litellm.compression.retrieval_tool import build_retrieval_tool
INPUT_TYPE = "openai_chat_completions"
ANTHROPIC_INPUT_TYPE = "anthropic_messages"
# ---------------------------------------------------------------------------
# BM25 scorer
@ -149,7 +152,7 @@ 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=INPUT_TYPE)
assert result["messages"] == messages
assert result["cache"] == {}
assert result["tools"] == []
@ -178,6 +181,7 @@ def test_compress_above_trigger():
result = litellm.compress(
big_messages,
model="gpt-4o",
input_type=INPUT_TYPE,
compression_trigger=1000,
compression_target=500,
)
@ -189,13 +193,62 @@ def test_compress_above_trigger():
assert result["tools"][0]["function"]["name"] == "litellm_content_retrieve"
def test_compress_anthropic_list_content_is_boundary_stable():
messages = [
{"role": "system", "content": [{"type": "text", "text": "System prompt"}]},
{
"role": "user",
"content": [
{"type": "text", "text": "# a.py\n" + "alpha " * 2000},
{
"type": "image_url",
"image_url": {"url": "https://example.com/a.png"},
},
],
},
{
"role": "user",
"content": [
{"type": "text", "text": "# b.py\n" + "beta " * 2000},
{
"type": "image_url",
"image_url": {"url": "https://example.com/b.png"},
},
],
},
{
"role": "user",
"content": [{"type": "text", "text": "Fix alpha bug in a.py"}],
},
]
result = litellm.compress(
messages=messages,
model="claude-sonnet-4-20250514",
input_type=ANTHROPIC_INPUT_TYPE,
compression_trigger=1000,
compression_target=500,
)
assert result["compressed_tokens"] < result["original_tokens"]
assert len(result["messages"]) == len(messages)
assert [m["role"] for m in result["messages"]] == [m["role"] for m in messages]
assert len(result["cache"]) > 0
assert len(result["tools"]) == 1
assert result["tools"][0]["type"] == "custom"
assert result["tools"][0]["name"] == "litellm_content_retrieve"
assert "input_schema" in result["tools"][0]
def test_compress_preserves_system_message():
messages = [
{"role": "system", "content": "System prompt. " * 500},
{"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=INPUT_TYPE, compression_trigger=1000
)
assert result["messages"][0]["role"] == "system"
assert "System prompt" in result["messages"][0]["content"]
@ -205,7 +258,9 @@ 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=INPUT_TYPE, 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 +271,9 @@ 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=INPUT_TYPE, 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 +286,9 @@ 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=INPUT_TYPE, compression_trigger=1000
)
if result["tools"]:
tool_desc = result["tools"][0]["function"]["description"]
for key in result["cache"]:
@ -242,11 +301,75 @@ def test_compress_default_target():
{"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=INPUT_TYPE, compression_trigger=2000
)
# Should have compressed — target = 1000
assert result["compressed_tokens"] <= result["original_tokens"]
def test_compress_nested_tool_result_extracts_text_only():
messages = [
{"role": "system", "content": [{"type": "text", "text": "System rules"}]},
{
"role": "user",
"content": [
{"type": "text", "text": "prefix"},
{
"type": "tool_result",
"tool_use_id": "toolu_1",
"content": [
{"type": "text", "text": "nested text fragment"},
{
"type": "image_url",
"image_url": {
"url": "https://example.com/secret-tool.png",
},
},
],
},
{
"type": "image_url",
"image_url": {"url": "https://example.com/top.png"},
},
{"type": "text", "text": " " + ("irrelevant " * 3000)},
],
},
{
"role": "user",
"content": [{"type": "text", "text": "final query that must remain"}],
},
]
result = litellm.compress(
messages=messages,
model="claude-sonnet-4-20250514",
input_type=ANTHROPIC_INPUT_TYPE,
compression_trigger=500,
compression_target=100,
)
cached_text = " ".join(result["cache"].values())
assert "nested text fragment" in cached_text
assert "https://example.com/secret-tool.png" not in cached_text
assert "https://example.com/top.png" not in cached_text
def test_compress_default_input_type_is_openai_chat_completions():
result = litellm.compress(
messages=[
{"role": "user", "content": "Large context " * 4000},
{"role": "user", "content": "query"},
],
model="gpt-4o",
compression_trigger=1000,
compression_target=500,
)
assert result["compressed_tokens"] <= result["original_tokens"]
assert isinstance(result["tools"], list)
def test_compress_forwards_embedding_model_params(monkeypatch):
captured = {}
@ -269,6 +392,7 @@ def test_compress_forwards_embedding_model_params(monkeypatch):
{"role": "user", "content": "Fix auth"},
],
model="gpt-4o",
input_type=INPUT_TYPE,
compression_trigger=1000,
embedding_model="text-embedding-3-small",
embedding_model_params={"api_base": "https://example-embeddings.test"},
@ -326,6 +450,7 @@ def test_embedding_scorer():
{"role": "user", "content": "Fix auth"},
],
model="gpt-4o",
input_type=INPUT_TYPE,
compression_trigger=1000,
embedding_model="text-embedding-3-small",
)
@ -346,8 +471,9 @@ 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)
print(result["messages"])
result = litellm.compress(
messages, model="gpt-4o", input_type=INPUT_TYPE, compression_trigger=1000
)
if expected_content == "Unrelated cooking recipes ":
assert "Unrelated cooking recipes " in result["messages"][1]["content"]
assert "Authentication code " not in result["messages"][0]["content"]