mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix: compress - make it work on anthropic input as well
This commit is contained in:
parent
9c20b8c743
commit
d1b9036dbf
15 changed files with 1050 additions and 93 deletions
|
|
@ -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).
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
14
litellm/integrations/compression_interception/__init__.py
Normal file
14
litellm/integrations/compression_interception/__init__.py
Normal 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",
|
||||
]
|
||||
382
litellm/integrations/compression_interception/handler.py
Normal file
382
litellm/integrations/compression_interception/handler.py
Normal 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
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
27
litellm/types/integrations/compression_interception.py
Normal file
27
litellm/types/integrations/compression_interception.py
Normal 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]]
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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]"
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue