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