Prompt Compression - add it to the proxy (#25729)

* refactor: new agentic loop event hook

simplifies how to create logic for tool based multi llm calls

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

* fix(compress.py): working prompt compression for claude code

ensures claude code messages can run through proxy easily

* docs: add agentic loop hook guide

* docs: add agentic_loop_hook to sidebar

* fix: fix multiple arguments error

* fix: fix tool call loop for compression on streaming /v1/messages

* fix: fix linting errors

* fix: fix ci/cd errors

* feat(litellm_pre_call_utils.py): use claude code session for litellm session id

allows claude code logs to be stitched together, making it easy to know they were all part of the same conversation

* fix: suppress incorrect mypy warning rE: module

* revert: drop PR's changes to litellm/proxy/_experimental/out/

Restores the 34 HTML files under _experimental/out/ to their pre-PR
paths (X/index.html -> X.html). All renames are R100 (content
unchanged); no other files are touched.

* fix: address greptile review comments on PR #25729

- Skip ``kwargs["tools"] = []`` injection when compression is a no-op —
  Anthropic Messages rejects empty tool arrays on requests that did not
  originally declare tools.
- Move agentic-loop safety guards (fingerprint cycle / max depth) out of
  the per-callback try/except so they propagate instead of being swallowed
  by the generic exception handler. Extracted _check_agentic_loop_safety.
- Gate generic ``x-<vendor>-session-id`` capture behind the
  LITELLM_CAPTURE_VENDOR_SESSION_HEADERS env var (off by default) to
  preserve backwards compatibility; explicit x-litellm-* headers are
  unaffected.
- Fix monkeypatch target in pre-call-hook test to patch the actual
  module-level binding
  (litellm.integrations.compression_interception.handler.compress).
- Add regression tests for empty-tools skip and opt-in session capture.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>

* revert: drop LITELLM_CAPTURE_VENDOR_SESSION_HEADERS flag

Generic x-<vendor>-session-id header capture is a new feature and only
runs *after* the explicit x-litellm-trace-id / x-litellm-session-id
checks, so it does not change behavior for any existing caller that was
already using the LiteLLM headers — no backwards-incompatibility to gate.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>

* refactor(compress): replace input_type with CallTypes call_type

Drop the bespoke ``CompressionInputType`` literal and use the existing
``litellm.types.utils.CallTypes`` enum instead.  ``litellm.compress()``
now takes ``call_type: Union[CallTypes, str]`` (default
``CallTypes.completion``) — no new concept to learn, and the enum is
already the way the rest of the codebase talks about request shapes.

Supported values: ``completion`` / ``acompletion`` (OpenAI chat-completions
shape) and ``anthropic_messages`` (Anthropic structured content blocks).

Updated: compress(), the compression_interception handler, tests, docs,
and the two eval scripts.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>

---------

Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
Krrish Dholakia
2026-04-20 15:08:00 -07:00
committed by GitHub
co-authored by Claude Opus 4.6
parent a19bff4ca6
commit 386f334fee
26 changed files with 3588 additions and 302 deletions
@@ -8,6 +8,7 @@ The function keeps high-relevance and recent context, replaces low-relevance con
```python
import litellm
from litellm.types.utils import CallTypes
messages = [
{"role": "system", "content": "You are a coding assistant."},
@@ -19,6 +20,7 @@ messages = [
compressed = litellm.compress(
messages=messages,
model="gpt-4o",
call_type=CallTypes.completion,
compression_trigger=1000,
compression_target=500,
)
@@ -45,6 +47,7 @@ response = litellm.completion(
- `messages` (`List[dict]`, required): input conversation messages
- `model` (`str`, required): model name used for token counting
- `call_type` (`CallTypes`, default `CallTypes.completion`): the LiteLLM call type whose message schema these messages follow. Supported values: `CallTypes.completion` / `CallTypes.acompletion` (OpenAI chat-completions shape) and `CallTypes.anthropic_messages` (Anthropic Messages shape)
- `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 +73,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).
@@ -0,0 +1,95 @@
# Agentic Loop Hook
Build a `CustomLogger` callback that intercepts a model response, fulfills tool calls server-side, and reruns the model — transparently to the caller.
:::info Supported call types
- `async` only (sync calls do not trigger the hook)
- Non-streaming only (streaming responses cannot be inspected for tool calls)
- Works on both `/v1/messages` and `/v1/chat/completions`
:::
## Implement the callback
Override two methods on `CustomLogger`:
```python
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.integrations.custom_logger import AgenticLoopPlan, AgenticLoopRequestPatch
MY_TOOL = "my_tool"
class MyToolCallback(CustomLogger):
async def async_should_run_agentic_loop(
self, response, model, messages, tools, stream, custom_llm_provider, kwargs
):
# Return (True, context_dict) if there are tool calls to handle
content = getattr(response, "content", None) or []
calls = [b for b in content if isinstance(b, dict)
and b.get("type") == "tool_use" and b.get("name") == MY_TOOL]
if not calls:
return False, {}
return True, {"tool_calls": calls}
async def async_build_agentic_loop_plan(
self, tools, model, messages, response,
anthropic_messages_provider_config,
anthropic_messages_optional_request_params,
logging_obj, stream, kwargs,
):
calls = tools["tool_calls"]
results = [f"result for {c['input']}" for c in calls] # your logic here
follow_up = messages + [
{"role": "assistant", "content": [
{"type": "tool_use", "id": c["id"], "name": c["name"], "input": c["input"]}
for c in calls
]},
{"role": "user", "content": [
{"type": "tool_result", "tool_use_id": c["id"], "content": results[i]}
for i, c in enumerate(calls)
]},
]
return AgenticLoopPlan(
run_agentic_loop=True,
request_patch=AgenticLoopRequestPatch(messages=follow_up),
)
```
For `/v1/chat/completions`, override `async_build_chat_completion_agentic_loop_plan` instead — same idea, `optional_params` replaces `anthropic_messages_optional_request_params`.
## Register it
```python
import litellm
litellm.callbacks = [MyToolCallback()]
```
Or in `config.yaml`:
```yaml
litellm_settings:
callbacks: ["my_module.MyToolCallback"]
```
## `AgenticLoopPlan` fields
| Field | Effect |
|---|---|
| `run_agentic_loop=True` + `request_patch` | Reruns the model with the patched request |
| `response_override` | Returns this value directly to the caller (no rerun) |
| `terminate=True` | Stops the loop, returns the current response |
| `run_agentic_loop=False` (default) | Skips; next callback is checked |
`AgenticLoopRequestPatch` accepts: `model`, `messages`, `tools`, `max_tokens`, `optional_params`, `kwargs`.
## Loop safety
- Default max reruns: `3` — override per-request with `kwargs["max_agentic_loops"]`
- Identical tool-call fingerprints abort the loop automatically
- Current depth is in `kwargs["_agentic_loop_depth"]`
## Examples in this repo
- `litellm/integrations/compression_interception/handler.py`
- `litellm/integrations/websearch_interception/handler.py`
+1
View File
@@ -536,6 +536,7 @@ const sidebars = {
description: "Modify requests, responses, and more",
items: [
"proxy/call_hooks",
"proxy/agentic_loop_hook",
"proxy/rules",
]
},
+1
View File
@@ -148,6 +148,7 @@ _custom_logger_compatible_callbacks_literal = Literal[
"vantage",
"posthog",
"levo",
"compression_interception",
]
cold_storage_custom_logger: Optional[_custom_logger_compatible_callbacks_literal] = None
logged_real_time_event_types: Optional[Union[List[str], Literal["*"]]] = None
+334 -78
View File
@@ -1,9 +1,9 @@
"""
Main compress() function — orchestrates BM25/embedding scoring, message stubbing,
and retrieval tool injection.
Main compress() function — normalizes input messages, orchestrates BM25/embedding
scoring, message stubbing, and retrieval tool injection.
"""
from typing import Any, Dict, List, Optional, Set, Union, cast
from typing import Any, Dict, List, Optional, Set, Tuple, Union, cast
from litellm.caching.dual_cache import DualCache
from litellm.compression.message_stubbing import (
@@ -15,27 +15,196 @@ 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.utils import AllMessageValues, Message
from litellm.types.utils import CallTypes
# CallTypes that produce Anthropic-shaped messages (structured content blocks).
# Everything else is treated as OpenAI chat-completions shape.
_ANTHROPIC_CALL_TYPES = frozenset({CallTypes.anthropic_messages.value})
# CallTypes that are valid targets for compression. Compression operates on
# message-shaped inputs, so we only accept call types whose payload is a list
# of role/content messages.
_SUPPORTED_CALL_TYPES = frozenset(
{
CallTypes.completion.value,
CallTypes.acompletion.value,
CallTypes.anthropic_messages.value,
}
)
def _normalize_call_type(call_type: Union[CallTypes, str]) -> str:
"""Return the string value for a ``CallTypes`` enum or a raw string."""
if isinstance(call_type, CallTypes):
return call_type.value
return call_type
def _is_anthropic_call_type(call_type: str) -> bool:
return call_type in _ANTHROPIC_CALL_TYPES
def _build_retrieval_tools(keys: List[str], call_type: str) -> List[dict]:
"""
Build retrieval tool definitions in the target request schema.
- Chat-completions call types: keep OpenAI function-tool schema.
- Anthropic messages call type: remap to Anthropic's custom tool schema.
"""
if not keys:
return []
openai_tools = [build_retrieval_tool(keys)]
if not _is_anthropic_call_type(call_type):
return openai_tools
# 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)
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`, expand into nested `content` items.
- Ignore non-textual blocks (images/documents/tool metadata/thinking metadata).
Implemented iteratively (stack-based) to avoid unbounded recursion.
"""
parts: List[str] = []
stack: List[Any] = [content]
while stack:
item = stack.pop()
if isinstance(item, str):
parts.append(item)
elif isinstance(item, list):
# Push list items in reverse order so they are processed left-to-right.
for element in reversed(item):
stack.append(element)
elif isinstance(item, dict):
item_type = item.get("type")
if item_type == "text":
parts.append(str(item.get("text", "")))
elif item_type == "tool_result":
stack.append(item.get("content", ""))
return " ".join(parts)
def _normalize_messages_for_compression(
messages: List[dict],
call_type: str,
) -> Tuple[List[dict], List[dict]]:
"""
Normalize each original message to a text-surrogate content for scoring.
Returns:
(normalized_messages, original_messages_copy)
"""
if call_type not in _SUPPORTED_CALL_TYPES:
raise ValueError(
f"Unsupported call_type={call_type!r} for compression. "
f"Expected one of: {sorted(_SUPPORTED_CALL_TYPES)}."
)
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 ""
def _extract_tool_use_ids(content: Any) -> List[str]:
if not isinstance(content, list):
return []
tool_use_ids: List[str] = []
for part in content:
if not isinstance(part, dict):
continue
if part.get("type") != "tool_use":
continue
tool_use_id = part.get("id")
if isinstance(tool_use_id, str) and tool_use_id:
tool_use_ids.append(tool_use_id)
return tool_use_ids
def _extract_tool_result_ids(content: Any) -> Set[str]:
if not isinstance(content, list):
return set()
tool_result_ids: Set[str] = set()
for part in content:
if not isinstance(part, dict):
continue
if part.get("type") != "tool_result":
continue
tool_use_id = part.get("tool_use_id")
if isinstance(tool_use_id, str) and tool_use_id:
tool_result_ids.add(tool_use_id)
return tool_result_ids
def _extract_anthropic_tool_exchange_spans(
messages: List[dict],
) -> Tuple[List[Set[int]], Optional[str]]:
"""
Return atomic 2-message spans for Anthropic tool exchanges.
Each assistant message containing `tool_use` must be immediately followed by a
user message containing matching `tool_result` blocks for all tool_use ids.
"""
spans: List[Set[int]] = []
i = 0
while i < len(messages):
current = messages[i]
if current.get("role") != "assistant":
i += 1
continue
tool_use_ids = _extract_tool_use_ids(current.get("content"))
if not tool_use_ids:
i += 1
continue
if i + 1 >= len(messages):
return [], "invalid_anthropic_tool_sequence"
next_msg = messages[i + 1]
if next_msg.get("role") != "user":
return [], "invalid_anthropic_tool_sequence"
tool_result_ids = _extract_tool_result_ids(next_msg.get("content"))
if not tool_result_ids:
return [], "invalid_anthropic_tool_sequence"
for tool_use_id in tool_use_ids:
if tool_use_id not in tool_result_ids:
return [], "invalid_anthropic_tool_sequence"
spans.append({i, i + 1})
i += 2
return spans, None
def _get_protected_indices(messages: List[dict]) -> List[int]:
"""
Return indices of messages that must never be compressed:
@@ -87,9 +256,98 @@ def _combine_scores(
return [bm25_weight * b + emb_weight * e for b, e in zip(norm_bm25, norm_emb)]
def _select_kept_indices_for_budget(
normalized_messages: List[dict],
original_messages: List[dict],
combined_scores: List[float],
compression_target: int,
model: str,
initial_kept_indices: Set[int],
tool_exchange_spans: List[Set[int]],
) -> Tuple[Set[int], Dict[int, dict]]:
kept_indices = set(initial_kept_indices)
current_tokens = 0
for i in kept_indices:
current_tokens += token_counter(
model=model,
text=cast(str, normalized_messages[i].get("content", "") or ""),
)
# Fill token budget from highest-scoring units.
# A unit is either:
# 1) a single message index, or
# 2) an Anthropic tool-exchange span that must be kept/dropped atomically.
truncated_overrides: Dict[int, dict] = {} # idx -> truncated message dict
span_id_by_index: Dict[int, int] = {}
for span_id, span in enumerate(tool_exchange_spans):
for idx in span:
span_id_by_index[idx] = span_id
# Build single-message candidate units (non-span messages).
candidate_units: List[Tuple[float, Tuple[int, ...], bool]] = []
for idx in range(len(normalized_messages)):
if idx in span_id_by_index or idx in kept_indices:
continue
candidate_units.append((combined_scores[idx], (idx,), True))
# Build span candidate units (atomic keep/drop for tool exchanges).
for span in tool_exchange_spans:
span_indices = tuple(sorted(span))
if any(idx in kept_indices for idx in span_indices):
continue
span_score = max(combined_scores[idx] for idx in span_indices)
candidate_units.append((span_score, span_indices, False))
# Sort by descending relevance score.
candidate_units.sort(key=lambda item: item[0], reverse=True)
for _score, indices, can_truncate in candidate_units:
if any(idx in kept_indices for idx in indices):
continue
msg_tokens = 0
for idx in indices:
msg_tokens += token_counter(
model=model,
text=cast(str, normalized_messages[idx].get("content", "") or ""),
)
remaining = compression_target - current_tokens
if remaining <= 0:
break # budget exhausted
if current_tokens + msg_tokens <= compression_target:
# Fits entirely
kept_indices.update(indices)
current_tokens += msg_tokens
elif can_truncate and len(indices) == 1 and remaining >= 100:
# Too large to fit whole single message, but we have budget — truncate it.
idx = indices[0]
truncated = truncate_message(original_messages[idx], remaining)
truncated_tokens = token_counter(
model=model,
text=truncated.get("content", "") or "",
)
truncated_overrides[idx] = truncated
kept_indices.add(idx)
current_tokens += truncated_tokens
return kept_indices, truncated_overrides
def _get_dropped_tool_span_indices(
kept_indices: Set[int], tool_exchange_spans: List[Set[int]]
) -> Set[int]:
dropped_tool_span_indices: Set[int] = set()
for span in tool_exchange_spans:
if not any(idx in kept_indices for idx in span):
dropped_tool_span_indices.update(span)
return dropped_tool_span_indices
def compress(
messages: List[dict],
model: str,
call_type: Union[CallTypes, str] = CallTypes.completion,
compression_trigger: int = 200_000,
compression_target: Optional[int] = None,
embedding_model: Optional[str] = None,
@@ -108,6 +366,12 @@ def compress(
Parameters:
messages: The conversation messages to (potentially) compress.
model: The LLM model name — used for token counting.
call_type: The LiteLLM call type whose message schema these messages
follow. Supported values:
- ``CallTypes.completion`` / ``CallTypes.acompletion`` — OpenAI
chat-completions shape (default)
- ``CallTypes.anthropic_messages`` — Anthropic Messages shape
(structured content blocks + atomic tool exchanges)
compression_trigger: Only compress if input exceeds this token count.
compression_target: Target token count after compression.
Defaults to ``compression_trigger // 2``.
@@ -122,29 +386,37 @@ def compress(
A ``CompressedResult`` dict containing compressed messages, token
counts, a cache of original content, and the retrieval tool definition.
"""
call_type_str = _normalize_call_type(call_type)
normalized_messages, original_messages = _normalize_messages_for_compression(
messages=messages,
call_type=call_type_str,
)
if compression_target is None:
compression_target = compression_trigger * 7 // 10
original_tokens = token_counter(
model=model, messages=cast(List[Union[AllMessageValues, Message]], messages)
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,
cache={},
tools=[],
compression_skipped_reason="below_trigger",
)
# 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 (
@@ -153,7 +425,7 @@ def compress(
emb_scores = embedding_score_messages(
query,
messages,
normalized_messages,
model=embedding_model,
cache=compression_cache,
embedding_model_params=embedding_model_params,
@@ -162,85 +434,69 @@ def compress(
else:
combined_scores = bm25_scores
# Sort message indices by score descending
ranked_indices = sorted(
range(len(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 ""
tool_exchange_spans: List[Set[int]] = []
if _is_anthropic_call_type(call_type_str):
tool_exchange_spans, tool_sequence_error = (
_extract_anthropic_tool_exchange_spans(original_messages)
)
# 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_tokens = token_counter(model=model, text=msg_content)
remaining = compression_target - current_tokens
if remaining <= 0:
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_tokens = token_counter(
model=model,
text=truncated.get("content", "") or "",
if tool_sequence_error is not None:
return CompressedResult(
messages=original_messages,
original_tokens=original_tokens,
compressed_tokens=original_tokens,
compression_ratio=0.0,
cache={},
tools=[],
compression_skipped_reason=tool_sequence_error,
)
truncated_overrides[idx] = truncated
kept_indices.add(idx)
current_tokens += truncated_tokens
for span in tool_exchange_spans:
# If any message in the span is protected, keep the whole span.
if any(idx in kept_indices for idx in span):
kept_indices.update(span)
kept_indices, truncated_overrides = _select_kept_indices_for_budget(
normalized_messages=normalized_messages,
original_messages=original_messages,
combined_scores=combined_scores,
compression_target=compression_target,
model=model,
initial_kept_indices=kept_indices,
tool_exchange_spans=tool_exchange_spans,
)
# Build compressed messages and cache
compressed_messages: List[dict] = []
cache: Dict[str, str] = {}
used_keys: Set[str] = set()
dropped_tool_span_indices = _get_dropped_tool_span_indices(
kept_indices=kept_indices, tool_exchange_spans=tool_exchange_spans
)
for i, msg in enumerate(messages):
for i, msg in enumerate(original_messages):
if i in dropped_tool_span_indices:
continue
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()), call_type=call_type_str)
compressed_tokens = token_counter(
model=model,
messages=cast(List[Union[AllMessageValues, Message]], compressed_messages),
messages=cast(List[Any], compressed_messages),
)
return CompressedResult(
@@ -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",
]
@@ -0,0 +1,399 @@
"""
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
from litellm._logging import verbose_logger
from litellm.compression import compress
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 = compress( # type: ignore
messages=messages,
model=model,
call_type=CallTypes.anthropic_messages,
compression_trigger=self.compression_trigger,
compression_target=self.compression_target,
embedding_model=self.embedding_model,
embedding_model_params=self.embedding_model_params,
)
cache = cast(Dict[str, str], compressed.get("cache", {}))
skip_reason = cast(Optional[str], compressed.get("compression_skipped_reason"))
compressed_tools = cast(List[Dict[str, Any]], compressed.get("tools", []))
# Only mutate kwargs when compression actually produced a result.
# If compression was a no-op (below trigger, invalid tool sequence, etc.),
# leave ``messages`` and ``tools`` untouched — injecting an empty
# ``tools: []`` onto a request that originally had no tools breaks
# Anthropic Messages requests.
if cache:
kwargs["messages"] = compressed["messages"]
if compressed_tools:
kwargs["tools"] = self._merge_tools(
existing_tools=cast(
Optional[List[Dict[str, Any]]], kwargs.get("tools")
),
compressed_tools=compressed_tools,
)
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),
)
elif skip_reason is not None:
verbose_logger.debug(
"CompressionInterception: compression skipped [reason=%s original=%d compressed=%d]",
skip_reason,
compressed.get("original_tokens"),
compressed.get("compressed_tokens"),
)
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
+37
View File
@@ -20,6 +20,7 @@ from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER
from litellm.types.integrations.argilla import ArgillaItem
from litellm.types.llms.openai import AllMessageValues, ChatCompletionRequest
from litellm.types.prompts.init_prompts import PromptSpec
from litellm.types.integrations.custom_logger import AgenticLoopPlan
from litellm.types.utils import (
AdapterCompletionStreamWrapper,
CallTypes,
@@ -676,6 +677,26 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
"""
pass
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: "LiteLLMLoggingObj",
stream: bool,
kwargs: Dict,
) -> AgenticLoopPlan:
"""
Build a typed rerun plan for Anthropic Messages agentic loops.
Override this method to separate callback decision/tool execution from
follow-up request execution (handled by BaseLLMHTTPHandler).
"""
return AgenticLoopPlan(run_agentic_loop=False)
async def async_should_run_chat_completion_agentic_loop(
self,
response: Any,
@@ -707,6 +728,22 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
"""
pass
async def async_build_chat_completion_agentic_loop_plan(
self,
tools: Dict,
model: str,
messages: List[Dict],
response: Any,
optional_params: Dict,
logging_obj: "LiteLLMLoggingObj",
stream: bool,
kwargs: Dict,
) -> AgenticLoopPlan:
"""
Build a typed rerun plan for chat-completions agentic loops.
"""
return AgenticLoopPlan(run_agentic_loop=False)
# Useful helpers for custom logger classes
def truncate_standard_logging_payload_content(
@@ -28,6 +28,10 @@ from litellm.integrations.websearch_interception.transformation import (
from litellm.types.integrations.websearch_interception import (
WebSearchInterceptionConfig,
)
from litellm.types.integrations.custom_logger import (
AgenticLoopPlan,
AgenticLoopRequestPatch,
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager
@@ -573,6 +577,35 @@ class WebSearchInterceptionLogger(CustomLogger):
kwargs=kwargs,
)
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:
tool_calls = tools["tool_calls"]
thinking_blocks = tools.get("thinking_blocks", [])
request_patch = await self._build_anthropic_request_patch(
model=model,
messages=messages,
tool_calls=tool_calls,
thinking_blocks=thinking_blocks,
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
logging_obj=logging_obj,
kwargs=kwargs,
)
return AgenticLoopPlan(
run_agentic_loop=True,
request_patch=request_patch,
metadata={"tool_type": "websearch", "response_format": "anthropic"},
)
async def async_run_chat_completion_agentic_loop(
self,
tools: Dict,
@@ -608,6 +641,33 @@ class WebSearchInterceptionLogger(CustomLogger):
response_format=response_format,
)
async def async_build_chat_completion_agentic_loop_plan(
self,
tools: Dict,
model: str,
messages: List[Dict],
response: Any,
optional_params: Dict,
logging_obj: Any,
stream: bool,
kwargs: Dict,
) -> AgenticLoopPlan:
tool_calls = tools["tool_calls"]
response_format = tools.get("response_format", "openai")
request_patch = await self._build_chat_completion_request_patch(
model=model,
messages=messages,
tool_calls=tool_calls,
optional_params=optional_params,
kwargs=kwargs,
response_format=response_format,
)
return AgenticLoopPlan(
run_agentic_loop=True,
request_patch=request_patch,
metadata={"tool_type": "websearch", "response_format": response_format},
)
@staticmethod
def _resolve_max_tokens(
optional_params: Dict,
@@ -672,7 +732,48 @@ class WebSearchInterceptionLogger(CustomLogger):
stream: bool,
kwargs: Dict,
) -> Any:
"""Execute litellm.search() and make follow-up request"""
"""Legacy path: execute search + build patch + run follow-up call."""
request_patch = await self._build_anthropic_request_patch(
model=model,
messages=messages,
tool_calls=tool_calls,
thinking_blocks=thinking_blocks,
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
logging_obj=logging_obj,
kwargs=kwargs,
)
if request_patch.messages is None:
raise ValueError("WebSearchInterception: missing follow-up messages")
optional_params = dict(anthropic_messages_optional_request_params)
optional_params.update(request_patch.optional_params)
max_tokens = request_patch.max_tokens
if max_tokens is None:
max_tokens = cast(Optional[int], optional_params.pop("max_tokens", None))
else:
optional_params.pop("max_tokens", None)
if max_tokens is None:
max_tokens = cast(int, kwargs.get("max_tokens", 1024))
return await anthropic_messages.acreate(
max_tokens=max_tokens,
messages=request_patch.messages,
model=request_patch.model or model,
**optional_params,
**request_patch.kwargs,
)
async def _build_anthropic_request_patch(
self,
model: str,
messages: List[Dict],
tool_calls: List[Dict],
thinking_blocks: List[Dict],
anthropic_messages_optional_request_params: Dict,
logging_obj: Any,
kwargs: Dict,
) -> AgenticLoopRequestPatch:
"""Execute litellm.search() and build follow-up request patch."""
# Extract search queries from tool_use blocks
search_tasks = []
@@ -721,20 +822,8 @@ class WebSearchInterceptionLogger(CustomLogger):
thinking_blocks=thinking_blocks,
)
# Make follow-up request with search results
# Type cast: user_message is a Dict for Anthropic format (default response_format)
follow_up_messages = messages + [assistant_message, cast(Dict, user_message)]
verbose_logger.debug(
"WebSearchInterception: Making follow-up request with search results"
)
verbose_logger.debug(
f"WebSearchInterception: Follow-up messages count: {len(follow_up_messages)}"
)
verbose_logger.debug(
f"WebSearchInterception: Last message (tool_result): {user_message}"
)
# Correlation context for structured logging
_call_id = getattr(logging_obj, "litellm_call_id", None) or kwargs.get(
"litellm_call_id", "unknown"
@@ -742,61 +831,39 @@ class WebSearchInterceptionLogger(CustomLogger):
full_model_name = model # safe default before try block
# Use anthropic_messages.acreate for follow-up request
try:
max_tokens = self._resolve_max_tokens(
anthropic_messages_optional_request_params, kwargs
)
max_tokens = self._resolve_max_tokens(
anthropic_messages_optional_request_params, kwargs
)
verbose_logger.debug(
f"WebSearchInterception: Using max_tokens={max_tokens} for follow-up request"
)
verbose_logger.debug(
f"WebSearchInterception: Using max_tokens={max_tokens} for follow-up request"
)
# Create a copy of optional params without max_tokens (since we pass it explicitly)
optional_params_without_max_tokens = {
k: v
for k, v in anthropic_messages_optional_request_params.items()
if k != "max_tokens"
}
optional_params_without_max_tokens = {
k: v
for k, v in anthropic_messages_optional_request_params.items()
if k != "max_tokens"
}
kwargs_for_followup = self._prepare_followup_kwargs(kwargs)
kwargs_for_followup = self._prepare_followup_kwargs(kwargs)
# Get model from logging_obj.model_call_details["agentic_loop_params"]
# This preserves the full model name with provider prefix (e.g., "bedrock/invoke/...")
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)
verbose_logger.debug(
f"WebSearchInterception: Using model name: {full_model_name}"
)
final_response = await anthropic_messages.acreate(
max_tokens=max_tokens,
messages=follow_up_messages,
model=full_model_name,
**optional_params_without_max_tokens,
**kwargs_for_followup,
)
verbose_logger.debug(
f"WebSearchInterception: Follow-up request completed, response type: {type(final_response)}"
)
verbose_logger.debug(
f"WebSearchInterception: Final response: {final_response}"
)
return final_response
except Exception as e:
verbose_logger.exception(
"WebSearchInterception: Follow-up request failed "
"[call_id=%s model=%s messages=%d searches=%d]: %s",
_call_id,
full_model_name,
len(follow_up_messages),
len(final_search_results),
str(e),
)
raise
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)
verbose_logger.debug(
"WebSearchInterception: Built anthropic request patch "
"[call_id=%s model=%s messages=%d searches=%d]",
_call_id,
full_model_name,
len(follow_up_messages),
len(final_search_results),
)
return AgenticLoopRequestPatch(
model=full_model_name,
messages=follow_up_messages,
max_tokens=max_tokens,
optional_params=optional_params_without_max_tokens,
kwargs=kwargs_for_followup,
)
async def _execute_search(self, query: str) -> str:
"""Execute a single web search using router's search tools"""
@@ -883,7 +950,36 @@ class WebSearchInterceptionLogger(CustomLogger):
kwargs: Dict,
response_format: str = "openai",
) -> Any:
"""Execute litellm.search() and make follow-up chat completion request"""
"""Legacy path: execute search + build patch + run follow-up call."""
request_patch = await self._build_chat_completion_request_patch(
model=model,
messages=messages,
tool_calls=tool_calls,
optional_params=optional_params,
kwargs=kwargs,
response_format=response_format,
)
if request_patch.messages is None:
raise ValueError("WebSearchInterception: missing follow-up messages")
params = dict(optional_params)
params.update(request_patch.optional_params)
return await litellm.acompletion(
model=request_patch.model or model,
messages=request_patch.messages,
**params,
**request_patch.kwargs,
)
async def _build_chat_completion_request_patch( # noqa: PLR0915
self,
model: str,
messages: List[Dict],
tool_calls: List[Dict],
optional_params: Dict,
kwargs: Dict,
response_format: str = "openai",
) -> AgenticLoopRequestPatch:
"""Execute litellm.search() and build chat-completion rerun patch."""
# Extract search queries from tool_calls
search_tasks = []
@@ -963,74 +1059,56 @@ class WebSearchInterceptionLogger(CustomLogger):
f"WebSearchInterception: Follow-up messages count: {len(follow_up_messages)}"
)
# Use litellm.acompletion for follow-up request
try:
# Remove internal parameters that shouldn't be passed to follow-up request
internal_params = {
"_websearch_interception",
"acompletion",
"litellm_logging_obj",
"custom_llm_provider",
# Remove internal parameters that shouldn't be passed to follow-up request
internal_params = {
"_websearch_interception",
"acompletion",
"litellm_logging_obj",
"custom_llm_provider",
"model_alias_map",
"stream_response",
"custom_prompt_dict",
}
kwargs_for_followup = {
k: v
for k, v in kwargs.items()
if not k.startswith("_websearch_interception") and k not in internal_params
}
full_model_name = model
if "custom_llm_provider" in kwargs:
custom_llm_provider = kwargs["custom_llm_provider"]
if not model.startswith(custom_llm_provider) and "/" not in model:
full_model_name = f"{custom_llm_provider}/{model}"
verbose_logger.debug(
"WebSearchInterception: Built chat completion request patch model=%s messages=%d",
full_model_name,
len(follow_up_messages),
)
tools_param = optional_params.get("tools")
optional_params_clean = {
k: v
for k, v in optional_params.items()
if k
not in {
"tools",
"extra_body",
"model_alias_map",
"stream_response",
"custom_prompt_dict",
}
kwargs_for_followup = {
k: v
for k, v in kwargs.items()
if not k.startswith("_websearch_interception")
and k not in internal_params
}
}
if tools_param is not None:
optional_params_clean["tools"] = tools_param
# Get full model name from kwargs
full_model_name = model
if "custom_llm_provider" in kwargs:
custom_llm_provider = kwargs["custom_llm_provider"]
# Reconstruct full model name with provider prefix if needed
if not model.startswith(custom_llm_provider):
# Check if model already has a provider prefix
if "/" not in model:
full_model_name = f"{custom_llm_provider}/{model}"
verbose_logger.debug(
f"WebSearchInterception: Using model name: {full_model_name}"
)
# Prepare tools for follow-up request (same as original)
tools_param = optional_params.get("tools")
# Remove tools and extra_body from optional_params to avoid issues
# extra_body often contains internal LiteLLM params that shouldn't be forwarded
optional_params_clean = {
k: v
for k, v in optional_params.items()
if k
not in {
"tools",
"extra_body",
"model_alias_map",
"stream_response",
"custom_prompt_dict",
}
}
final_response = await litellm.acompletion(
model=full_model_name,
messages=follow_up_messages,
tools=tools_param,
**optional_params_clean,
**kwargs_for_followup,
)
verbose_logger.debug(
f"WebSearchInterception: Follow-up request completed, response type: {type(final_response)}"
)
return final_response
except Exception as e:
verbose_logger.exception(
f"WebSearchInterception: Follow-up request failed: {str(e)}"
)
raise
return AgenticLoopRequestPatch(
model=full_model_name,
messages=follow_up_messages,
optional_params=optional_params_clean,
kwargs=kwargs_for_followup,
)
async def _create_empty_search_result(self) -> str:
"""Create an empty search result for tool calls without queries"""
@@ -0,0 +1,320 @@
"""
Agentic Streaming Iterator for Anthropic Messages
Wraps the raw SSE byte stream from the Anthropic pass-through endpoint,
yields every chunk to the caller (preserving real streaming), collects
all bytes, and on stream exhaustion rebuilds the full Anthropic response
to run through agentic completion hooks. If an agentic hook fires, the
follow-up response is chained as Phase 2 of the same iterator.
"""
import json
from typing import Any, AsyncIterator, Dict, List, Optional, cast
from litellm._logging import verbose_logger
# ---------------------------------------------------------------------------
# SSE parsing helpers (module-level to keep the class lean)
# ---------------------------------------------------------------------------
def _parse_sse_events(raw: bytes) -> List[tuple]:
"""Return a list of (event_type, parsed_data_dict) from raw SSE bytes."""
text = raw.decode("utf-8", errors="replace")
lines = text.split("\n")
events: List[tuple] = []
current_event_type: Optional[str] = None
for line in lines:
stripped = line.strip()
if stripped.startswith("event:"):
current_event_type = stripped[len("event:") :].strip()
continue
if not stripped.startswith("data:"):
continue
data_str = stripped[len("data:") :].strip()
try:
data = json.loads(data_str)
except (json.JSONDecodeError, ValueError):
continue
event_type = current_event_type or data.get("type", "")
current_event_type = None
events.append((event_type, data))
return events
def _handle_message_start(data: Dict, response: Dict) -> None:
msg = data.get("message", {})
response["id"] = msg.get("id", response["id"])
response["model"] = msg.get("model", response["model"])
response["role"] = msg.get("role", response["role"])
usage = msg.get("usage", {})
if usage:
response["usage"]["input_tokens"] = usage.get("input_tokens", 0)
for key in ("cache_creation_input_tokens", "cache_read_input_tokens"):
if key in usage:
response["usage"][key] = usage[key]
def _handle_content_block_start(data: Dict, content_blocks: Dict[int, Dict]) -> None:
idx = data.get("index", len(content_blocks))
block = data.get("content_block", {})
block_type = block.get("type", "text")
_BLOCK_TEMPLATES: Dict[str, Dict] = {
"text": {"type": "text", "text": ""},
"thinking": {"type": "thinking", "thinking": "", "signature": ""},
"redacted_thinking": {
"type": "redacted_thinking",
"data": block.get("data", ""),
},
}
if block_type == "tool_use":
content_blocks[idx] = {
"type": "tool_use",
"id": block.get("id", ""),
"name": block.get("name", ""),
"input": {},
"_partial_json": "",
}
elif block_type in _BLOCK_TEMPLATES:
content_blocks[idx] = dict(_BLOCK_TEMPLATES[block_type])
else:
content_blocks[idx] = dict(block)
def _handle_content_block_delta(data: Dict, content_blocks: Dict[int, Dict]) -> None:
idx = data.get("index", 0)
delta = data.get("delta", {})
delta_type = delta.get("type", "")
block = content_blocks.get(idx)
if block is None:
return
if delta_type == "text_delta":
block["text"] = block.get("text", "") + delta.get("text", "")
elif delta_type == "input_json_delta":
block["_partial_json"] = block.get("_partial_json", "") + delta.get(
"partial_json", ""
)
elif delta_type == "thinking_delta":
block["thinking"] = block.get("thinking", "") + delta.get("thinking", "")
elif delta_type == "signature_delta":
block["signature"] = delta.get("signature", block.get("signature", ""))
def _handle_content_block_stop(data: Dict, content_blocks: Dict[int, Dict]) -> None:
idx = data.get("index", 0)
block = content_blocks.get(idx)
if block and block.get("type") == "tool_use":
partial = block.pop("_partial_json", "")
if partial:
try:
block["input"] = json.loads(partial)
except (json.JSONDecodeError, ValueError):
block["input"] = {"_raw": partial}
def _handle_message_delta(data: Dict, response: Dict) -> None:
delta = data.get("delta", {})
if "stop_reason" in delta:
response["stop_reason"] = delta["stop_reason"]
if "stop_sequence" in delta:
response["stop_sequence"] = delta["stop_sequence"]
usage = data.get("usage", {})
if usage.get("output_tokens") is not None:
response["usage"]["output_tokens"] = usage["output_tokens"]
for key in (
"input_tokens",
"cache_creation_input_tokens",
"cache_read_input_tokens",
):
if key in usage:
response["usage"][key] = usage[key]
class AgenticAnthropicStreamingIterator:
"""
Two-phase async iterator that enables agentic hooks on streaming
Anthropic Messages pass-through responses.
Phase 1: Yield raw SSE bytes from the upstream response while
accumulating them. When the inner iterator is exhausted,
rebuild the full Anthropic response dict and call agentic hooks.
Phase 2: If an agentic hook fires and returns a follow-up response
(streaming or non-streaming), yield those bytes to the caller.
"""
def __init__(
self,
completion_stream: AsyncIterator,
http_handler: Any,
model: str,
messages: List[Dict],
anthropic_messages_provider_config: Any,
anthropic_messages_optional_request_params: Dict,
logging_obj: Any,
custom_llm_provider: str,
kwargs: Dict,
):
self._inner = completion_stream.__aiter__()
self._http_handler = http_handler
self._model = model
self._messages = messages
self._anthropic_messages_provider_config = anthropic_messages_provider_config
self._anthropic_messages_optional_request_params = (
anthropic_messages_optional_request_params
)
self._logging_obj = logging_obj
self._custom_llm_provider = custom_llm_provider
self._kwargs = kwargs
self._collected_bytes: List[bytes] = []
self._stream_exhausted = False
self._hook_processing_done = False
self._follow_up_iterator: Optional[AsyncIterator] = None
def __aiter__(self):
return self
async def __anext__(self) -> bytes:
# Phase 1: yield from upstream, collect bytes
if not self._stream_exhausted:
try:
chunk = await self._inner.__anext__()
self._collected_bytes.append(chunk)
return chunk
except StopAsyncIteration:
self._stream_exhausted = True
await self._process_agentic_hooks()
# Fall through to Phase 2
# Phase 2: yield from follow-up stream if one was created
if self._follow_up_iterator is not None:
chunk = await self._follow_up_iterator.__anext__()
return chunk
raise StopAsyncIteration
async def _process_agentic_hooks(self) -> None:
"""Rebuild the Anthropic response from collected SSE bytes and call hooks."""
if self._hook_processing_done:
return
self._hook_processing_done = True
if not self._collected_bytes:
return
try:
rebuilt = self._rebuild_anthropic_response_from_sse(self._collected_bytes)
if rebuilt is None:
verbose_logger.debug(
"AgenticStreamingIterator: Could not rebuild response from SSE bytes"
)
return
[
f"{b.get('type')}({b.get('name', '')})"
if b.get("type") == "tool_use"
else b.get("type")
for b in rebuilt.get("content", [])
]
result = await self._http_handler._call_agentic_completion_hooks(
response=rebuilt,
model=self._model,
messages=self._messages,
anthropic_messages_provider_config=self._anthropic_messages_provider_config,
anthropic_messages_optional_request_params=self._anthropic_messages_optional_request_params,
logging_obj=self._logging_obj,
stream=True,
custom_llm_provider=self._custom_llm_provider,
kwargs=self._kwargs,
)
if result is None:
return
if hasattr(result, "__aiter__"):
self._follow_up_iterator = result.__aiter__()
elif isinstance(result, dict):
from litellm.llms.anthropic.experimental_pass_through.messages.fake_stream_iterator import (
FakeAnthropicMessagesStreamIterator,
)
from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
)
fake = FakeAnthropicMessagesStreamIterator(
response=cast(AnthropicMessagesResponse, result)
)
self._follow_up_iterator = fake.__aiter__()
else:
verbose_logger.warning(
"AgenticStreamingIterator: Unexpected result type from hooks: %s",
type(result).__name__,
)
except Exception as e:
_call_id = getattr(self._logging_obj, "litellm_call_id", "unknown")
verbose_logger.exception(
"AgenticStreamingIterator: Error in agentic hook processing "
"[call_id=%s model=%s]: %s",
_call_id,
self._model,
str(e),
)
@staticmethod
def _rebuild_anthropic_response_from_sse(
raw_bytes: List[bytes],
) -> Optional[Dict[str, Any]]:
"""
Parse collected SSE bytes into an Anthropic Messages response dict.
Processes SSE events in order:
- message_start -> envelope (id, model, role, usage)
- content_block_start -> new content block
- content_block_delta -> accumulate text/json/thinking deltas
- content_block_stop -> finalize block
- message_delta -> stop_reason, output usage
- message_stop -> end
"""
events = _parse_sse_events(b"".join(raw_bytes))
response: Dict[str, Any] = {
"id": "",
"type": "message",
"role": "assistant",
"model": "",
"content": [],
"stop_reason": None,
"stop_sequence": None,
"usage": {"input_tokens": 0, "output_tokens": 0},
}
content_blocks: Dict[int, Dict[str, Any]] = {}
saw_message_start = False
for event_type, data in events:
if event_type == "message_start":
saw_message_start = True
_handle_message_start(data, response)
elif event_type == "content_block_start":
_handle_content_block_start(data, content_blocks)
elif event_type == "content_block_delta":
_handle_content_block_delta(data, content_blocks)
elif event_type == "content_block_stop":
_handle_content_block_stop(data, content_blocks)
elif event_type == "message_delta":
_handle_message_delta(data, response)
if not saw_message_start:
return None
for idx in sorted(content_blocks.keys()):
block = content_blocks[idx]
block.pop("_partial_json", None)
response["content"].append(block)
return response
+369 -70
View File
@@ -78,6 +78,10 @@ from litellm.types.containers.main import (
DeleteContainerResult,
)
from litellm.types.files import TwoStepFileUploadConfig
from litellm.types.integrations.custom_logger import (
AgenticLoopPlan,
AgenticLoopRequestPatch,
)
from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
)
@@ -2047,7 +2051,23 @@ class BaseLLMHTTPHandler:
request_body=request_body,
litellm_logging_obj=logging_obj,
)
initial_response = completion_stream
from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import (
AgenticAnthropicStreamingIterator,
)
initial_response = AgenticAnthropicStreamingIterator(
completion_stream=completion_stream,
http_handler=self,
model=model,
messages=messages,
anthropic_messages_provider_config=anthropic_messages_provider_config,
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
kwargs=kwargs,
)
return initial_response
else:
initial_response = anthropic_messages_provider_config.transform_anthropic_messages_response(
model=model,
@@ -2055,7 +2075,7 @@ class BaseLLMHTTPHandler:
logging_obj=logging_obj,
)
# Call agentic completion hooks
# Call agentic completion hooks (non-streaming path only)
final_response = await self._call_agentic_completion_hooks(
response=initial_response,
model=model,
@@ -2063,7 +2083,7 @@ class BaseLLMHTTPHandler:
anthropic_messages_provider_config=anthropic_messages_provider_config,
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
logging_obj=logging_obj,
stream=stream or False,
stream=False,
custom_llm_provider=custom_llm_provider,
kwargs=kwargs,
)
@@ -4516,6 +4536,167 @@ class BaseLLMHTTPHandler:
return stream, data
return stream, data
@staticmethod
def _get_agentic_loop_settings(kwargs: Dict) -> Tuple[int, int, List[str]]:
depth = int(kwargs.get("_agentic_loop_depth", 0) or 0)
max_loops = int(kwargs.get("max_agentic_loops", 3) or 3)
fingerprints = list(kwargs.get("_agentic_loop_fingerprints", []) or [])
return depth, max(max_loops, 1), fingerprints
@staticmethod
def _check_agentic_loop_safety(
tool_calls: Any,
fingerprints: List[str],
depth: int,
max_loops: int,
model: str,
) -> str:
"""
Evaluate agentic-loop safety guards (fingerprint cycle / max depth).
Raises ValueError on abort. Returns the current fingerprint on success.
These checks must not be swallowed by the per-callback ``except Exception``
block that wraps callback dispatch they are bounded-loop / cycle-break
safety rails and must abort the agentic dispatch when they trip.
"""
fingerprint = BaseLLMHTTPHandler._fingerprint_agentic_tools(tool_calls)
if fingerprint in fingerprints:
raise ValueError(
"Agentic loop detected repeated tool-call fingerprint; aborting rerun"
)
if depth >= max_loops:
raise ValueError(
f"Exceeded max_agentic_loops={max_loops} for model={model}"
)
return fingerprint
@staticmethod
def _fingerprint_agentic_tools(tools: Dict) -> str:
try:
return json.dumps(tools, sort_keys=True, default=str)
except Exception:
return str(tools)
async def _execute_anthropic_agentic_plan(
self,
plan: AgenticLoopPlan,
model: str,
messages: List[Dict],
anthropic_messages_optional_request_params: Dict,
logging_obj: "LiteLLMLoggingObj",
kwargs: Dict,
depth: int,
max_loops: int,
fingerprints: List[str],
fingerprint: str,
stream: bool = False,
) -> Any:
from litellm.anthropic_interface import messages as anthropic_messages
patch = plan.request_patch or AgenticLoopRequestPatch()
if patch.messages is None:
raise ValueError("Agentic loop plan missing patched messages")
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))
optional_params = dict(anthropic_messages_optional_request_params)
optional_params.update(patch.optional_params)
if patch.tools is not None:
optional_params["tools"] = patch.tools
max_tokens = patch.max_tokens
if max_tokens is None:
max_tokens = cast(Optional[int], optional_params.pop("max_tokens", None))
else:
optional_params.pop("max_tokens", None)
if max_tokens is None:
max_tokens = cast(int, kwargs.get("max_tokens", 1024))
internal_keys = {"litellm_logging_obj"}
kwargs_for_followup = {
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_keys
and k not in optional_params
}
kwargs_for_followup.update(patch.kwargs)
kwargs_for_followup["_agentic_loop_depth"] = depth + 1
kwargs_for_followup["max_agentic_loops"] = max_loops
kwargs_for_followup["_agentic_loop_fingerprints"] = fingerprints + [fingerprint]
return await anthropic_messages.acreate(
**{
"max_tokens": max_tokens,
"messages": patch.messages,
"model": patch.model or full_model_name,
"stream": stream,
**optional_params,
**kwargs_for_followup,
}
)
async def _execute_chat_completion_agentic_plan(
self,
plan: AgenticLoopPlan,
model: str,
messages: List[Dict],
optional_params: Dict,
kwargs: Dict,
custom_llm_provider: str,
depth: int,
max_loops: int,
fingerprints: List[str],
fingerprint: str,
) -> Any:
patch = plan.request_patch or AgenticLoopRequestPatch()
if patch.messages is None:
raise ValueError("Agentic loop plan missing patched messages")
full_model_name = patch.model or model
if "/" not in full_model_name:
full_model_name = f"{custom_llm_provider}/{full_model_name}"
optional_params_for_followup = dict(optional_params)
optional_params_for_followup.update(patch.optional_params)
if patch.tools is not None:
optional_params_for_followup["tools"] = patch.tools
internal_params = {
"_websearch_interception",
"acompletion",
"litellm_logging_obj",
"custom_llm_provider",
"model_alias_map",
"stream_response",
"custom_prompt_dict",
}
kwargs_for_followup = {
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)
kwargs_for_followup["_agentic_loop_depth"] = depth + 1
kwargs_for_followup["max_agentic_loops"] = max_loops
kwargs_for_followup["_agentic_loop_fingerprints"] = fingerprints + [fingerprint]
return await litellm.acompletion(
model=full_model_name,
messages=patch.messages,
**optional_params_for_followup,
**kwargs_for_followup,
)
async def _call_agentic_completion_hooks(
self,
response: Any,
@@ -4541,45 +4722,111 @@ class BaseLLMHTTPHandler:
callbacks = litellm.callbacks + (logging_obj.dynamic_success_callbacks or [])
tools = anthropic_messages_optional_request_params.get("tools", [])
depth, max_loops, fingerprints = self._get_agentic_loop_settings(kwargs=kwargs)
for callback in callbacks:
if not isinstance(callback, CustomLogger):
continue
should_run: bool = False
tool_calls: Any = None
try:
if isinstance(callback, CustomLogger):
# First: Check if agentic loop should run
(
should_run,
tool_calls,
) = await callback.async_should_run_agentic_loop(
response=response,
# First: Check if agentic loop should run. Wrap in try/except
# to shield from buggy user callbacks — a callback crash should
# not abort the whole request.
(
should_run,
tool_calls,
) = await callback.async_should_run_agentic_loop(
response=response,
model=model,
messages=messages,
tools=tools,
stream=stream,
custom_llm_provider=custom_llm_provider,
kwargs=kwargs,
)
except Exception as e:
_call_id = getattr(logging_obj, "litellm_call_id", "unknown")
verbose_logger.exception(
"LiteLLM.AgenticHookError: Exception in "
"async_should_run_agentic_loop [call_id=%s model=%s]: %s",
_call_id,
model,
str(e),
)
continue
if not should_run:
continue
# Safety guards must run OUTSIDE the callback try/except — they are
# bounded-loop / cycle-break rails that must propagate to the caller.
fingerprint = self._check_agentic_loop_safety(
tool_calls=tool_calls,
fingerprints=fingerprints,
depth=depth,
max_loops=max_loops,
model=model,
)
try:
kwargs_with_provider = kwargs.copy() if kwargs else {}
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
)
if not build_plan_overridden:
return await callback.async_run_agentic_loop(
tools=tool_calls,
model=model,
messages=messages,
tools=tools,
response=response,
anthropic_messages_provider_config=anthropic_messages_provider_config,
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
logging_obj=logging_obj,
stream=stream,
custom_llm_provider=custom_llm_provider,
kwargs=kwargs,
kwargs=kwargs_with_provider,
)
if should_run:
# Second: Execute agentic loop
# Add custom_llm_provider to kwargs so the agentic loop can reconstruct the full model name
kwargs_with_provider = kwargs.copy() if kwargs else {}
kwargs_with_provider["custom_llm_provider"] = (
custom_llm_provider
)
agentic_response = await callback.async_run_agentic_loop(
tools=tool_calls,
model=model,
messages=messages,
response=response,
anthropic_messages_provider_config=anthropic_messages_provider_config,
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
logging_obj=logging_obj,
stream=stream,
kwargs=kwargs_with_provider,
)
# First hook that runs agentic loop wins
return agentic_response
plan = await callback.async_build_agentic_loop_plan(
tools=tool_calls,
model=model,
messages=messages,
response=response,
anthropic_messages_provider_config=anthropic_messages_provider_config,
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
logging_obj=logging_obj,
stream=stream,
kwargs=kwargs_with_provider,
)
if plan.response_override is not None:
return plan.response_override
if plan.terminate:
verbose_logger.debug(
"Agentic loop terminated by callback=%s reason=%s",
callback.__class__.__name__,
plan.stop_reason,
)
return response
if not plan.run_agentic_loop:
continue
return await self._execute_anthropic_agentic_plan(
plan=plan,
model=model,
messages=messages,
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
logging_obj=logging_obj,
kwargs=kwargs_with_provider,
depth=depth,
max_loops=max_loops,
fingerprints=fingerprints,
fingerprint=fingerprint,
stream=stream,
)
except Exception as e:
_call_id = getattr(logging_obj, "litellm_call_id", "unknown")
verbose_logger.exception(
@@ -4653,52 +4900,104 @@ class BaseLLMHTTPHandler:
callbacks = litellm.callbacks + (logging_obj.dynamic_success_callbacks or [])
tools = optional_params.get("tools", [])
depth, max_loops, fingerprints = self._get_agentic_loop_settings(kwargs=kwargs)
for callback in callbacks:
try:
if isinstance(callback, CustomLogger):
# Check if callback has the chat completion agentic loop method
if not hasattr(
callback, "async_should_run_chat_completion_agentic_loop"
):
continue
if not isinstance(callback, CustomLogger):
continue
if not hasattr(callback, "async_should_run_chat_completion_agentic_loop"):
continue
# First: Check if agentic loop should run
(
should_run,
tool_calls,
) = await callback.async_should_run_chat_completion_agentic_loop(
response=response,
should_run: bool = False
tool_calls: Any = None
try:
(
should_run,
tool_calls,
) = await callback.async_should_run_chat_completion_agentic_loop(
response=response,
model=model,
messages=messages,
tools=tools,
stream=stream,
custom_llm_provider=custom_llm_provider,
kwargs=kwargs,
)
except Exception as e:
verbose_logger.exception(
"LiteLLM.AgenticHookError: Exception in "
"async_should_run_chat_completion_agentic_loop: %s",
str(e),
)
continue
if not should_run:
continue
# Safety guards must run OUTSIDE the callback try/except — they are
# bounded-loop / cycle-break rails that must propagate to the caller.
fingerprint = self._check_agentic_loop_safety(
tool_calls=tool_calls,
fingerprints=fingerprints,
depth=depth,
max_loops=max_loops,
model=model,
)
try:
kwargs_with_provider = kwargs.copy() if kwargs else {}
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,
tools=tools,
response=response,
optional_params=optional_params,
logging_obj=logging_obj,
stream=stream,
custom_llm_provider=custom_llm_provider,
kwargs=kwargs,
kwargs=kwargs_with_provider,
)
if should_run:
# Second: Execute agentic loop
# Add custom_llm_provider to kwargs so the agentic loop can reconstruct the full model name
kwargs_with_provider = kwargs.copy() if kwargs else {}
kwargs_with_provider["custom_llm_provider"] = (
custom_llm_provider
)
agentic_response = (
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,
)
)
# First hook that runs agentic loop wins
return agentic_response
plan = await callback.async_build_chat_completion_agentic_loop_plan(
tools=tool_calls,
model=model,
messages=messages,
response=response,
optional_params=optional_params,
logging_obj=logging_obj,
stream=stream,
kwargs=kwargs_with_provider,
)
if plan.response_override is not None:
return plan.response_override
if plan.terminate:
verbose_logger.debug(
"Agentic chat loop terminated by callback=%s reason=%s",
callback.__class__.__name__,
plan.stop_reason,
)
return response
if not plan.run_agentic_loop:
continue
return await self._execute_chat_completion_agentic_plan(
plan=plan,
model=model,
messages=messages,
optional_params=optional_params,
kwargs=kwargs_with_provider,
custom_llm_provider=custom_llm_provider,
depth=depth,
max_loops=max_loops,
fingerprints=fingerprints,
fingerprint=fingerprint,
)
except Exception as e:
verbose_logger.exception(
f"LiteLLM.AgenticHookError: Exception in chat completion agentic hooks: {str(e)}"
+13 -3
View File
@@ -22,11 +22,21 @@ model_list:
output_cost_per_token: 10 # 100x standard ($10.00/1M = $0.00001)
# Anthropic model for /v1/messages test — 100x custom pricing
- model_name: "claude-sonnet-4-20250514"
- model_name: "claude-sonnet-4-6"
litellm_params:
model: anthropic/claude-sonnet-4-20250514
model: anthropic/claude-sonnet-4-6
api_key: os.environ/ANTHROPIC_API_KEY
model_info:
id: claude-sonnet-4-custom-pricing
input_cost_per_token: 0.0003 # 100x standard ($0.000003)
output_cost_per_token: 0.0015 # 100x standard ($0.000015)
output_cost_per_token: 0.0015 # 100x standard ($0.000015)
litellm_settings:
callbacks: ["compression_interception"]
compression_interception_params:
enabled: true
compression_trigger: 100000
# # 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(
+48 -9
View File
@@ -1,5 +1,6 @@
import asyncio
import copy
import re
import time
from collections import OrderedDict
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
@@ -28,6 +29,14 @@ _SPECIAL_HEADERS_CACHE = frozenset(
v.value.lower() for v in SpecialHeaders._member_map_.values()
)
# Matches any header of the form x-<something>-session-id (case-insensitive).
# Excludes the two explicit litellm headers which are handled with higher priority.
_GENERIC_SESSION_ID_HEADER_RE = re.compile(r"^x-.+-session-id$", re.IGNORECASE)
_EXPLICIT_SESSION_HEADERS = frozenset({"x-litellm-trace-id", "x-litellm-session-id"})
# Session-id values must be non-empty strings of alphanumerics, hyphens, or underscores
# (covers UUIDs and most common session-id formats).
_SESSION_ID_VALUE_RE = re.compile(r"^[a-zA-Z0-9_\-]{8,}$")
def _sanitize_for_log(value: Any) -> str:
"""
@@ -115,13 +124,43 @@ def _get_metadata_variable_name(request: Request) -> str:
return "metadata"
def _extract_generic_session_id_from_headers(
normalized: Dict[str, str],
) -> Optional[str]:
"""
Scan a normalised (lower-cased keys) header dict for any header that looks
like ``x-<vendor>-session-id`` and whose value is a plausible session/trace
identifier (alphanumeric + hyphens/underscores, at least 8 chars).
The two explicit LiteLLM headers (``x-litellm-trace-id`` /
``x-litellm-session-id``) are excluded here because they are handled with
higher priority by the caller.
Example: ``x-claude-code-session-id: e96634a3-fa28-4083-b354-55542e2dca01``
"""
for key, value in normalized.items():
if (
key not in _EXPLICIT_SESSION_HEADERS
and _GENERIC_SESSION_ID_HEADER_RE.match(key)
and isinstance(value, str)
and _SESSION_ID_VALUE_RE.match(value)
):
return value
return None
def get_chain_id_from_headers(headers: Optional[Dict[str, str]]) -> Optional[str]:
"""
Extract chain id for call chaining from request headers.
x-litellm-trace-id and x-litellm-session-id are interchangeable; when both
are present, x-litellm-trace-id takes precedence. Header keys are matched
case-insensitively so this works with raw header dicts from any transport.
Priority order:
1. ``x-litellm-trace-id`` (explicit, highest priority)
2. ``x-litellm-session-id`` (explicit)
3. Any ``x-<vendor>-session-id`` header whose value looks like a session id
(alphanumeric / UUID, at least 8 chars). E.g. ``x-claude-code-session-id``.
Header keys are matched case-insensitively so this works with raw header
dicts from any transport.
Used by MCP (and other paths that have raw_headers but no Request) to set
litellm_trace_id/litellm_session_id for spend logs and logging consistency.
@@ -129,8 +168,10 @@ def get_chain_id_from_headers(headers: Optional[Dict[str, str]]) -> Optional[str
if not headers:
return None
normalized = {k.lower(): v for k, v in headers.items() if isinstance(k, str)}
return normalized.get("x-litellm-trace-id") or normalized.get(
"x-litellm-session-id"
return (
normalized.get("x-litellm-trace-id")
or normalized.get("x-litellm-session-id")
or _extract_generic_session_id_from_headers(normalized)
)
@@ -649,10 +690,8 @@ class LiteLLMProxyRequestSetup:
#########################################################################################
agent_id_from_header = headers.get("x-litellm-agent-id")
# x-litellm-trace-id and x-litellm-session-id are interchangeable for call chaining
chain_id = headers.get("x-litellm-trace-id") or headers.get(
"x-litellm-session-id"
)
# Explicit litellm headers take precedence; fall back to any x-*-session-id header.
chain_id = get_chain_id_from_headers(dict(headers))
if agent_id_from_header:
metadata_from_headers["agent_id"] = agent_id_from_header
+9 -1
View File
@@ -2,7 +2,14 @@
Type definitions for litellm.compress().
"""
from typing import Dict, List, TypedDict
import sys
if sys.version_info >= (3, 11):
from typing import Dict, List, NotRequired, TypedDict
else:
from typing import Dict, List, TypedDict
from typing_extensions import NotRequired
class CompressedResult(TypedDict):
@@ -12,3 +19,4 @@ class CompressedResult(TypedDict):
compression_ratio: float # fraction reduced, e.g. 0.6 means 60% reduction
cache: Dict[str, str] # key -> original content (for retrieval tool responses)
tools: List[dict] # [litellm_content_retrieve tool definition]
compression_skipped_reason: NotRequired[str]
@@ -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]]
+28 -2
View File
@@ -1,6 +1,6 @@
from typing import Optional
from typing import Any, Dict, List, Optional
from pydantic import BaseModel
from pydantic import BaseModel, Field
class StandardCustomLoggerInitParams(BaseModel):
@@ -9,3 +9,29 @@ class StandardCustomLoggerInitParams(BaseModel):
"""
turn_off_message_logging: Optional[bool] = False
class AgenticLoopRequestPatch(BaseModel):
"""
Patch returned by callbacks to request a follow-up LLM call.
"""
model: Optional[str] = None
messages: Optional[List[Dict[str, Any]]] = None
tools: Optional[List[Dict[str, Any]]] = None
max_tokens: Optional[int] = None
optional_params: Dict[str, Any] = Field(default_factory=dict)
kwargs: Dict[str, Any] = Field(default_factory=dict)
class AgenticLoopPlan(BaseModel):
"""
Typed callback response for agentic-loop reruns.
"""
run_agentic_loop: bool = False
request_patch: Optional[AgenticLoopRequestPatch] = None
response_override: Optional[Any] = None
terminate: bool = False
stop_reason: Optional[str] = None
metadata: Dict[str, Any] = Field(default_factory=dict)
+2
View File
@@ -33,6 +33,7 @@ from dataclasses import asdict, dataclass, field
from typing import Optional
import litellm
from litellm.types.utils import CallTypes
# ---------------------------------------------------------------------------
# Problem definitions (HumanEval-style)
@@ -880,6 +881,7 @@ def eval_problem(
result = litellm.compress(
messages=messages,
model=model,
call_type=CallTypes.completion,
compression_trigger=compression_trigger,
embedding_model=embedding_model,
)
+2 -1
View File
@@ -40,6 +40,7 @@ sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
import litellm # noqa: E402
from litellm.compression import compress as litellm_compress # noqa: E402
from litellm.types.utils import CallTypes # noqa: E402
# ---------------------------------------------------------------------------
# Prompts
@@ -445,7 +446,7 @@ def eval_instance(
compress_kwargs: dict = {
"messages": messages,
"model": model,
"input_type": "openai_chat_completions",
"call_type": CallTypes.completion,
"compression_trigger": compression_trigger,
"embedding_model": embedding_model,
}
@@ -0,0 +1,364 @@
"""
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
# The handler does ``from litellm.compression import compress`` at module
# scope, so we must patch the binding on the handler module — patching
# ``litellm.compress`` has no effect on the already-bound reference.
monkeypatch.setattr(
"litellm.integrations.compression_interception.handler.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_pre_call_hook_below_trigger_does_not_inject_empty_tools(monkeypatch):
"""
When compression is a no-op (below trigger / invalid tool sequence), the
hook must NOT replace ``messages`` or inject an empty ``tools: []`` onto
a request that originally had no tools Anthropic Messages rejects
``tools: []``.
"""
logger = CompressionInterceptionLogger()
original_messages = [{"role": "user", "content": "short prompt"}]
def _fake_compress_noop(**kwargs):
return {
"messages": original_messages,
"original_tokens": 42,
"compressed_tokens": 42,
"compression_ratio": 0.0,
"cache": {},
"tools": [],
"compression_skipped_reason": "below_trigger",
}
monkeypatch.setattr(
"litellm.integrations.compression_interception.handler.compress",
_fake_compress_noop,
)
kwargs = {
"model": "bedrock/us.anthropic.claude-sonnet-4-5",
"messages": original_messages,
}
result = await logger.async_pre_call_deployment_hook(
kwargs=kwargs, call_type=CallTypes.anthropic_messages
)
assert result is not None
# Original request had no ``tools`` — skipped compression must leave it that way.
assert "tools" not in result
# Cache must not be populated for a no-op.
assert result.get("litellm_call_id") not 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_should_run_agentic_loop_with_custom_type_tools():
"""Test that async_should_run_agentic_loop returns True when tools contain
litellm_content_retrieve as a custom-typed tool (e.g. Claude Code tool list)
and the model response includes a matching tool_use block."""
logger = CompressionInterceptionLogger()
# Exact tools payload produced by Claude Code litellm_content_retrieve is
# the final entry and uses type="custom" (not type="function").
tools = [
{
"name": "Agent",
"description": "Launch a new agent to handle complex, multi-step tasks.",
"input_schema": {
"$schema": "https://json-schema.org/draft/2020-12/schema",
"type": "object",
"properties": {
"description": {"type": "string"},
"prompt": {"type": "string"},
},
"required": ["description", "prompt"],
"additionalProperties": False,
},
},
{
"name": "AskUserQuestion",
"description": "Use this tool when you need to ask the user questions.",
"input_schema": {
"$schema": "https://json-schema.org/draft/2020-12/schema",
"type": "object",
"properties": {
"questions": {"type": "array", "items": {"type": "object"}},
},
"required": ["questions"],
"additionalProperties": False,
},
},
{
"name": "Bash",
"description": "Executes a given bash command and returns its output.",
"input_schema": {
"$schema": "https://json-schema.org/draft/2020-12/schema",
"type": "object",
"properties": {"command": {"type": "string"}},
"required": ["command"],
"additionalProperties": False,
},
},
{
"name": "litellm_content_retrieve",
"description": "Retrieve the full content of a file or message that was compressed to save tokens.",
"input_schema": {
"type": "object",
"properties": {
"key": {
"type": "string",
"description": "The identifier of the content to retrieve",
"enum": [
"message_0",
"HA_UPTIME_ROUTER_SPEC.md",
"message_159",
"message_160",
],
}
},
"required": ["key"],
},
"type": "custom",
},
]
response = {
"content": [
{
"type": "tool_use",
"id": "toolu_abc",
"name": "litellm_content_retrieve",
"input": {"key": "message_0"},
}
]
}
should_run, tools_dict = await logger.async_should_run_agentic_loop(
response=response,
model="claude-3-5-sonnet",
messages=[],
tools=tools,
stream=False,
custom_llm_provider="anthropic",
kwargs={},
)
assert should_run is True
assert tools_dict["tool_type"] == "compression_retrieval"
assert len(tools_dict["tool_calls"]) == 1
assert tools_dict["tool_calls"][0]["input"]["key"] == "message_0"
@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]"
)
@@ -4,7 +4,7 @@ Unit tests for WebSearch Interception Handler
Tests the WebSearchInterceptionLogger class and helper functions.
"""
from unittest.mock import MagicMock, Mock
from unittest.mock import AsyncMock, MagicMock, Mock
import pytest
@@ -69,6 +69,61 @@ async def test_async_should_run_agentic_loop():
assert tools_dict == {}
@pytest.mark.asyncio
async def test_async_build_agentic_loop_plan_returns_request_patch():
"""Callback should return a typed patch for base handler reruns."""
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"])
logger._execute_search = AsyncMock( # type: ignore
return_value="Title: LiteLLM\nURL: docs\nSnippet: test"
)
tools_dict = {
"tool_calls": [
{
"id": "toolu_123",
"type": "tool_use",
"name": "litellm_web_search",
"input": {"query": "what is litellm"},
}
],
"response_format": "anthropic",
}
logging_obj = MagicMock()
logging_obj.model_call_details = {
"agentic_loop_params": {"model": "bedrock/invoke/claude-3-5-sonnet"}
}
kwargs = {
"temperature": 0.2,
"_websearch_interception_converted_stream": True,
"litellm_logging_obj": object(),
}
plan = await logger.async_build_agentic_loop_plan(
tools=tools_dict,
model="claude-3-5-sonnet",
messages=[{"role": "user", "content": "search LiteLLM"}],
response=None,
anthropic_messages_provider_config=None,
anthropic_messages_optional_request_params={
"max_tokens": 1024,
"tools": [{"name": "litellm_web_search"}],
},
logging_obj=logging_obj,
stream=False,
kwargs=kwargs,
)
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
assert "_websearch_interception_converted_stream" not in plan.request_patch.kwargs
assert "litellm_logging_obj" not in plan.request_patch.kwargs
assert plan.request_patch.kwargs["temperature"] == 0.2
@pytest.mark.asyncio
async def test_internal_flags_filtered_from_followup_kwargs():
"""Test that internal _websearch_interception flags are filtered from follow-up request kwargs.
@@ -0,0 +1,792 @@
"""
Tests for AgenticAnthropicStreamingIterator and SSE rebuild helpers.
"""
import json
import os
import sys
from typing import Any, Dict, List, Optional, Tuple
from unittest.mock import AsyncMock, MagicMock
import pytest
sys.path.insert(0, os.path.abspath("../../../../.."))
from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import (
AgenticAnthropicStreamingIterator,
_handle_content_block_delta,
_handle_content_block_start,
_handle_content_block_stop,
_handle_message_delta,
_handle_message_start,
_parse_sse_events,
)
# ---------------------------------------------------------------------------
# Helpers to build SSE byte payloads
# ---------------------------------------------------------------------------
def _sse_event(event_type: str, data: dict) -> bytes:
return f"event: {event_type}\ndata: {json.dumps(data)}\n\n".encode()
def _build_simple_text_stream() -> List[bytes]:
"""Produce SSE bytes for a simple text response (no tool calls)."""
chunks = []
chunks.append(
_sse_event(
"message_start",
{
"type": "message_start",
"message": {
"id": "msg_123",
"type": "message",
"role": "assistant",
"model": "claude-sonnet-4-20250514",
"content": [],
"stop_reason": None,
"stop_sequence": None,
"usage": {"input_tokens": 10, "output_tokens": 0},
},
},
)
)
chunks.append(
_sse_event(
"content_block_start",
{
"type": "content_block_start",
"index": 0,
"content_block": {"type": "text", "text": ""},
},
)
)
chunks.append(
_sse_event(
"content_block_delta",
{
"type": "content_block_delta",
"index": 0,
"delta": {"type": "text_delta", "text": "Hello, world!"},
},
)
)
chunks.append(
_sse_event("content_block_stop", {"type": "content_block_stop", "index": 0})
)
chunks.append(
_sse_event(
"message_delta",
{
"type": "message_delta",
"delta": {"stop_reason": "end_turn", "stop_sequence": None},
"usage": {"output_tokens": 5},
},
)
)
chunks.append(_sse_event("message_stop", {"type": "message_stop"}))
return chunks
def _build_tool_use_stream() -> List[bytes]:
"""Produce SSE bytes for a response with a tool_use block."""
chunks = []
chunks.append(
_sse_event(
"message_start",
{
"type": "message_start",
"message": {
"id": "msg_tool_456",
"type": "message",
"role": "assistant",
"model": "claude-sonnet-4-20250514",
"content": [],
"stop_reason": None,
"usage": {"input_tokens": 50, "output_tokens": 0},
},
},
)
)
# thinking block
chunks.append(
_sse_event(
"content_block_start",
{
"type": "content_block_start",
"index": 0,
"content_block": {
"type": "thinking",
"thinking": "",
"signature": "",
},
},
)
)
chunks.append(
_sse_event(
"content_block_delta",
{
"type": "content_block_delta",
"index": 0,
"delta": {
"type": "thinking_delta",
"thinking": "I need to retrieve...",
},
},
)
)
chunks.append(
_sse_event(
"content_block_delta",
{
"type": "content_block_delta",
"index": 0,
"delta": {"type": "signature_delta", "signature": "sig_abc"},
},
)
)
chunks.append(
_sse_event("content_block_stop", {"type": "content_block_stop", "index": 0})
)
# tool_use block
chunks.append(
_sse_event(
"content_block_start",
{
"type": "content_block_start",
"index": 1,
"content_block": {
"type": "tool_use",
"id": "toolu_001",
"name": "litellm_content_retrieve",
"input": {},
},
},
)
)
chunks.append(
_sse_event(
"content_block_delta",
{
"type": "content_block_delta",
"index": 1,
"delta": {
"type": "input_json_delta",
"partial_json": '{"key": "section_',
},
},
)
)
chunks.append(
_sse_event(
"content_block_delta",
{
"type": "content_block_delta",
"index": 1,
"delta": {"type": "input_json_delta", "partial_json": '1"}'},
},
)
)
chunks.append(
_sse_event("content_block_stop", {"type": "content_block_stop", "index": 1})
)
chunks.append(
_sse_event(
"message_delta",
{
"type": "message_delta",
"delta": {"stop_reason": "tool_use"},
"usage": {"output_tokens": 20},
},
)
)
chunks.append(_sse_event("message_stop", {"type": "message_stop"}))
return chunks
# ---------------------------------------------------------------------------
# Mock async stream
# ---------------------------------------------------------------------------
class MockAsyncStream:
"""Async iterator that yields a list of byte chunks."""
def __init__(self, chunks: List[bytes]):
self._chunks = list(chunks)
self._idx = 0
def __aiter__(self):
return self
async def __anext__(self) -> bytes:
if self._idx >= len(self._chunks):
raise StopAsyncIteration
chunk = self._chunks[self._idx]
self._idx += 1
return chunk
# ---------------------------------------------------------------------------
# Tests for _parse_sse_events
# ---------------------------------------------------------------------------
class TestParseSSEEvents:
def test_should_parse_single_event(self):
raw = _sse_event(
"message_start", {"type": "message_start", "message": {"id": "1"}}
)
events = _parse_sse_events(raw)
assert len(events) == 1
assert events[0][0] == "message_start"
assert events[0][1]["message"]["id"] == "1"
def test_should_parse_multiple_events(self):
raw = b"".join(_build_simple_text_stream())
events = _parse_sse_events(raw)
event_types = [e[0] for e in events]
assert "message_start" in event_types
assert "content_block_start" in event_types
assert "content_block_delta" in event_types
assert "content_block_stop" in event_types
assert "message_delta" in event_types
assert "message_stop" in event_types
def test_should_skip_malformed_json(self):
raw = b"event: message_start\ndata: {invalid json}\n\n"
events = _parse_sse_events(raw)
assert len(events) == 0
def test_should_handle_empty_bytes(self):
events = _parse_sse_events(b"")
assert events == []
# ---------------------------------------------------------------------------
# Tests for _handle_* helpers
# ---------------------------------------------------------------------------
class TestHandleMessageStart:
def test_should_populate_envelope(self):
response: Dict[str, Any] = {
"id": "",
"model": "",
"role": "assistant",
"usage": {"input_tokens": 0, "output_tokens": 0},
}
data = {
"message": {
"id": "msg_abc",
"model": "claude-sonnet-4-20250514",
"role": "assistant",
"usage": {
"input_tokens": 42,
"cache_creation_input_tokens": 100,
},
}
}
_handle_message_start(data, response)
assert response["id"] == "msg_abc"
assert response["model"] == "claude-sonnet-4-20250514"
assert response["usage"]["input_tokens"] == 42
assert response["usage"]["cache_creation_input_tokens"] == 100
class TestHandleContentBlockStart:
def test_should_create_text_block(self):
blocks: Dict[int, Dict] = {}
data = {"index": 0, "content_block": {"type": "text", "text": ""}}
_handle_content_block_start(data, blocks)
assert blocks[0] == {"type": "text", "text": ""}
def test_should_create_tool_use_block(self):
blocks: Dict[int, Dict] = {}
data = {
"index": 1,
"content_block": {
"type": "tool_use",
"id": "toolu_x",
"name": "my_tool",
"input": {},
},
}
_handle_content_block_start(data, blocks)
assert blocks[1]["type"] == "tool_use"
assert blocks[1]["name"] == "my_tool"
assert blocks[1]["_partial_json"] == ""
def test_should_create_thinking_block(self):
blocks: Dict[int, Dict] = {}
data = {
"index": 0,
"content_block": {"type": "thinking", "thinking": "", "signature": ""},
}
_handle_content_block_start(data, blocks)
assert blocks[0]["type"] == "thinking"
class TestHandleContentBlockDelta:
def test_should_accumulate_text(self):
blocks = {0: {"type": "text", "text": "Hello"}}
_handle_content_block_delta(
{"index": 0, "delta": {"type": "text_delta", "text": " World"}},
blocks,
)
assert blocks[0]["text"] == "Hello World"
def test_should_accumulate_json(self):
blocks = {0: {"type": "tool_use", "_partial_json": '{"key":'}}
_handle_content_block_delta(
{
"index": 0,
"delta": {"type": "input_json_delta", "partial_json": '"val"}'},
},
blocks,
)
assert blocks[0]["_partial_json"] == '{"key":"val"}'
def test_should_ignore_missing_block(self):
blocks: Dict[int, Dict] = {}
_handle_content_block_delta(
{"index": 99, "delta": {"type": "text_delta", "text": "x"}},
blocks,
)
assert 99 not in blocks
class TestHandleContentBlockStop:
def test_should_parse_tool_input_json(self):
blocks = {
0: {
"type": "tool_use",
"input": {},
"_partial_json": '{"key": "section_1"}',
}
}
_handle_content_block_stop({"index": 0}, blocks)
assert blocks[0]["input"] == {"key": "section_1"}
assert "_partial_json" not in blocks[0]
def test_should_handle_invalid_json_gracefully(self):
blocks = {
0: {
"type": "tool_use",
"input": {},
"_partial_json": "not valid json",
}
}
_handle_content_block_stop({"index": 0}, blocks)
assert blocks[0]["input"] == {"_raw": "not valid json"}
class TestHandleMessageDelta:
def test_should_set_stop_reason_and_usage(self):
response: Dict[str, Any] = {
"stop_reason": None,
"stop_sequence": None,
"usage": {"input_tokens": 0, "output_tokens": 0},
}
_handle_message_delta(
{
"delta": {"stop_reason": "end_turn", "stop_sequence": None},
"usage": {"output_tokens": 15},
},
response,
)
assert response["stop_reason"] == "end_turn"
assert response["usage"]["output_tokens"] == 15
# ---------------------------------------------------------------------------
# Tests for _rebuild_anthropic_response_from_sse
# ---------------------------------------------------------------------------
class TestRebuildAnthropicResponse:
def test_should_rebuild_simple_text_response(self):
raw_bytes = _build_simple_text_stream()
result = AgenticAnthropicStreamingIterator._rebuild_anthropic_response_from_sse(
raw_bytes
)
assert result is not None
assert result["id"] == "msg_123"
assert result["model"] == "claude-sonnet-4-20250514"
assert result["stop_reason"] == "end_turn"
assert len(result["content"]) == 1
assert result["content"][0]["type"] == "text"
assert result["content"][0]["text"] == "Hello, world!"
assert result["usage"]["input_tokens"] == 10
assert result["usage"]["output_tokens"] == 5
def test_should_rebuild_tool_use_response(self):
raw_bytes = _build_tool_use_stream()
result = AgenticAnthropicStreamingIterator._rebuild_anthropic_response_from_sse(
raw_bytes
)
assert result is not None
assert result["id"] == "msg_tool_456"
assert result["stop_reason"] == "tool_use"
assert len(result["content"]) == 2
thinking = result["content"][0]
assert thinking["type"] == "thinking"
assert thinking["thinking"] == "I need to retrieve..."
assert thinking["signature"] == "sig_abc"
tool = result["content"][1]
assert tool["type"] == "tool_use"
assert tool["id"] == "toolu_001"
assert tool["name"] == "litellm_content_retrieve"
assert tool["input"] == {"key": "section_1"}
def test_should_return_none_without_message_start(self):
raw_bytes = [
_sse_event(
"content_block_start",
{
"type": "content_block_start",
"index": 0,
"content_block": {"type": "text"},
},
)
]
result = AgenticAnthropicStreamingIterator._rebuild_anthropic_response_from_sse(
raw_bytes
)
assert result is None
def test_should_handle_empty_bytes(self):
result = AgenticAnthropicStreamingIterator._rebuild_anthropic_response_from_sse(
[]
)
assert result is None
def test_should_handle_multi_event_chunks(self):
"""When multiple SSE events arrive in a single bytes chunk."""
combined = b"".join(_build_simple_text_stream())
result = AgenticAnthropicStreamingIterator._rebuild_anthropic_response_from_sse(
[combined]
)
assert result is not None
assert result["content"][0]["text"] == "Hello, world!"
def test_should_preserve_cache_usage_fields(self):
raw_bytes = [
_sse_event(
"message_start",
{
"type": "message_start",
"message": {
"id": "msg_cache",
"model": "claude-sonnet-4-20250514",
"role": "assistant",
"usage": {
"input_tokens": 100,
"cache_creation_input_tokens": 50,
"cache_read_input_tokens": 30,
},
},
},
),
_sse_event(
"message_delta",
{
"type": "message_delta",
"delta": {"stop_reason": "end_turn"},
"usage": {"output_tokens": 10},
},
),
_sse_event("message_stop", {"type": "message_stop"}),
]
result = AgenticAnthropicStreamingIterator._rebuild_anthropic_response_from_sse(
raw_bytes
)
assert result is not None
assert result["usage"]["cache_creation_input_tokens"] == 50
assert result["usage"]["cache_read_input_tokens"] == 30
def test_should_handle_redacted_thinking_block(self):
raw_bytes = [
_sse_event(
"message_start",
{
"type": "message_start",
"message": {
"id": "msg_redact",
"model": "claude-sonnet-4-20250514",
"role": "assistant",
"usage": {"input_tokens": 5},
},
},
),
_sse_event(
"content_block_start",
{
"type": "content_block_start",
"index": 0,
"content_block": {"type": "redacted_thinking", "data": "abc123"},
},
),
_sse_event(
"content_block_stop",
{"type": "content_block_stop", "index": 0},
),
_sse_event(
"message_delta",
{
"type": "message_delta",
"delta": {"stop_reason": "end_turn"},
"usage": {"output_tokens": 1},
},
),
_sse_event("message_stop", {"type": "message_stop"}),
]
result = AgenticAnthropicStreamingIterator._rebuild_anthropic_response_from_sse(
raw_bytes
)
assert result is not None
assert result["content"][0]["type"] == "redacted_thinking"
# ---------------------------------------------------------------------------
# Tests for AgenticAnthropicStreamingIterator (Phase 1 / Phase 2)
# ---------------------------------------------------------------------------
class TestAgenticStreamingIteratorPhase1:
@pytest.mark.asyncio
async def test_should_yield_all_chunks_when_no_hook_fires(self):
"""When hooks return None, the wrapper should yield all original chunks."""
chunks = _build_simple_text_stream()
mock_stream = MockAsyncStream(chunks)
mock_handler = MagicMock()
mock_handler._call_agentic_completion_hooks = AsyncMock(return_value=None)
iterator = AgenticAnthropicStreamingIterator(
completion_stream=mock_stream,
http_handler=mock_handler,
model="claude-sonnet-4-20250514",
messages=[{"role": "user", "content": "hi"}],
anthropic_messages_provider_config=MagicMock(),
anthropic_messages_optional_request_params={},
logging_obj=MagicMock(),
custom_llm_provider="anthropic",
kwargs={},
)
collected = []
async for chunk in iterator:
collected.append(chunk)
assert len(collected) == len(chunks)
for orig, got in zip(chunks, collected):
assert orig == got
mock_handler._call_agentic_completion_hooks.assert_awaited_once()
@pytest.mark.asyncio
async def test_should_pass_rebuilt_response_to_hooks(self):
"""The rebuilt dict passed to hooks should match the original stream content."""
chunks = _build_tool_use_stream()
mock_stream = MockAsyncStream(chunks)
captured_response = {}
async def mock_hooks(**kwargs):
captured_response.update(kwargs["response"])
return None
mock_handler = MagicMock()
mock_handler._call_agentic_completion_hooks = mock_hooks
iterator = AgenticAnthropicStreamingIterator(
completion_stream=mock_stream,
http_handler=mock_handler,
model="claude-sonnet-4-20250514",
messages=[],
anthropic_messages_provider_config=MagicMock(),
anthropic_messages_optional_request_params={},
logging_obj=MagicMock(),
custom_llm_provider="anthropic",
kwargs={},
)
async for _ in iterator:
pass
assert captured_response["id"] == "msg_tool_456"
assert captured_response["stop_reason"] == "tool_use"
assert captured_response["content"][1]["name"] == "litellm_content_retrieve"
class TestAgenticStreamingIteratorPhase2:
@pytest.mark.asyncio
async def test_should_chain_follow_up_async_iterator(self):
"""When hooks return an async iterator, Phase 2 should yield from it."""
phase1_chunks = _build_simple_text_stream()
phase2_chunks = [b"follow-up-chunk-1", b"follow-up-chunk-2"]
mock_stream = MockAsyncStream(phase1_chunks)
follow_up = MockAsyncStream(phase2_chunks)
mock_handler = MagicMock()
mock_handler._call_agentic_completion_hooks = AsyncMock(return_value=follow_up)
iterator = AgenticAnthropicStreamingIterator(
completion_stream=mock_stream,
http_handler=mock_handler,
model="claude-sonnet-4-20250514",
messages=[],
anthropic_messages_provider_config=MagicMock(),
anthropic_messages_optional_request_params={},
logging_obj=MagicMock(),
custom_llm_provider="anthropic",
kwargs={},
)
collected = []
async for chunk in iterator:
collected.append(chunk)
assert len(collected) == len(phase1_chunks) + len(phase2_chunks)
assert collected[-2:] == phase2_chunks
@pytest.mark.asyncio
async def test_should_convert_dict_response_to_fake_stream(self):
"""When hooks return a dict, it should be wrapped in FakeAnthropicMessagesStreamIterator."""
phase1_chunks = _build_simple_text_stream()
mock_stream = MockAsyncStream(phase1_chunks)
fake_response = {
"id": "msg_followup",
"type": "message",
"role": "assistant",
"model": "claude-sonnet-4-20250514",
"content": [{"type": "text", "text": "follow-up answer"}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 100, "output_tokens": 20},
}
mock_handler = MagicMock()
mock_handler._call_agentic_completion_hooks = AsyncMock(
return_value=fake_response
)
iterator = AgenticAnthropicStreamingIterator(
completion_stream=mock_stream,
http_handler=mock_handler,
model="claude-sonnet-4-20250514",
messages=[],
anthropic_messages_provider_config=MagicMock(),
anthropic_messages_optional_request_params={},
logging_obj=MagicMock(),
custom_llm_provider="anthropic",
kwargs={},
)
collected = []
async for chunk in iterator:
collected.append(chunk)
# Phase 1 chunks + Phase 2 fake-stream chunks
assert len(collected) > len(phase1_chunks)
# The follow-up chunks should contain the text from the dict response
phase2_bytes = b"".join(collected[len(phase1_chunks) :])
assert b"follow-up answer" in phase2_bytes
class TestAgenticStreamingIteratorErrorHandling:
@pytest.mark.asyncio
async def test_should_swallow_hook_errors(self):
"""Errors in hook processing should be swallowed; Phase 1 chunks are still yielded."""
chunks = _build_simple_text_stream()
mock_stream = MockAsyncStream(chunks)
mock_handler = MagicMock()
mock_handler._call_agentic_completion_hooks = AsyncMock(
side_effect=RuntimeError("hook exploded")
)
mock_logging = MagicMock()
mock_logging.litellm_call_id = "test_call_123"
iterator = AgenticAnthropicStreamingIterator(
completion_stream=mock_stream,
http_handler=mock_handler,
model="claude-sonnet-4-20250514",
messages=[],
anthropic_messages_provider_config=MagicMock(),
anthropic_messages_optional_request_params={},
logging_obj=mock_logging,
custom_llm_provider="anthropic",
kwargs={},
)
collected = []
async for chunk in iterator:
collected.append(chunk)
# All Phase 1 chunks should still have been yielded
assert len(collected) == len(chunks)
@pytest.mark.asyncio
async def test_should_handle_empty_stream(self):
"""An empty upstream stream should not crash."""
mock_stream = MockAsyncStream([])
mock_handler = MagicMock()
mock_handler._call_agentic_completion_hooks = AsyncMock(return_value=None)
iterator = AgenticAnthropicStreamingIterator(
completion_stream=mock_stream,
http_handler=mock_handler,
model="claude-sonnet-4-20250514",
messages=[],
anthropic_messages_provider_config=MagicMock(),
anthropic_messages_optional_request_params={},
logging_obj=MagicMock(),
custom_llm_provider="anthropic",
kwargs={},
)
collected = []
async for chunk in iterator:
collected.append(chunk)
assert collected == []
# hooks should not be called since no bytes were collected
mock_handler._call_agentic_completion_hooks.assert_not_awaited()
@pytest.mark.asyncio
async def test_should_pass_stream_true_to_hooks(self):
"""The wrapper should always pass stream=True to hooks."""
chunks = _build_simple_text_stream()
mock_stream = MockAsyncStream(chunks)
mock_handler = MagicMock()
mock_handler._call_agentic_completion_hooks = AsyncMock(return_value=None)
iterator = AgenticAnthropicStreamingIterator(
completion_stream=mock_stream,
http_handler=mock_handler,
model="claude-sonnet-4-20250514",
messages=[],
anthropic_messages_provider_config=MagicMock(),
anthropic_messages_optional_request_params={},
logging_obj=MagicMock(),
custom_llm_provider="anthropic",
kwargs={},
)
async for _ in iterator:
pass
call_kwargs = mock_handler._call_agentic_completion_hooks.call_args
assert call_kwargs.kwargs["stream"] is True
@@ -74,6 +74,34 @@ def test_prepare_fake_stream_request():
assert result_data["messages"] == [{"role": "user", "content": "Hello"}]
def test_get_agentic_loop_settings_defaults_and_overrides():
handler = BaseLLMHTTPHandler()
depth, max_loops, fingerprints = handler._get_agentic_loop_settings(kwargs={})
assert depth == 0
assert max_loops == 3
assert fingerprints == []
depth, max_loops, fingerprints = handler._get_agentic_loop_settings(
kwargs={
"_agentic_loop_depth": 2,
"max_agentic_loops": 7,
"_agentic_loop_fingerprints": ["fp-1", "fp-2"],
}
)
assert depth == 2
assert max_loops == 7
assert fingerprints == ["fp-1", "fp-2"]
def test_fingerprint_agentic_tools_is_deterministic():
handler = BaseLLMHTTPHandler()
tools_a = {"tool_calls": [{"id": "1", "input": {"q": "abc"}, "name": "web_search"}]}
tools_b = {"tool_calls": [{"name": "web_search", "input": {"q": "abc"}, "id": "1"}]}
assert handler._fingerprint_agentic_tools(tools_a) == handler._fingerprint_agentic_tools(tools_b)
@pytest.mark.asyncio
async def test_async_anthropic_messages_handler_extra_headers():
"""
@@ -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
@@ -84,3 +87,35 @@ def test_normalize_callback_names_lowercases_strings():
"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
@@ -1791,6 +1791,57 @@ def test_add_litellm_metadata_from_request_headers_both_headers_trace_id_precede
assert data["litellm_trace_id"] == "trace-value"
def test_add_litellm_metadata_from_request_headers_generic_session_id_header():
"""A generic x-<vendor>-session-id header is used when no explicit litellm header is set."""
headers = {"x-claude-code-session-id": "e96634a3-fa28-4083-b354-55542e2dca01"}
data = {"metadata": {}}
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
headers=headers, data=data, _metadata_variable_name="metadata"
)
assert data["metadata"]["session_id"] == "e96634a3-fa28-4083-b354-55542e2dca01"
assert data["litellm_session_id"] == "e96634a3-fa28-4083-b354-55542e2dca01"
assert data["litellm_trace_id"] == "e96634a3-fa28-4083-b354-55542e2dca01"
def test_add_litellm_metadata_from_request_headers_explicit_header_beats_generic():
"""Explicit x-litellm-trace-id wins over a generic x-*-session-id header."""
headers = {
"x-litellm-trace-id": "explicit-trace-id-value",
"x-claude-code-session-id": "e96634a3-fa28-4083-b354-55542e2dca01",
}
data = {"metadata": {}}
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
headers=headers, data=data, _metadata_variable_name="metadata"
)
assert data["litellm_session_id"] == "explicit-trace-id-value"
assert data["litellm_trace_id"] == "explicit-trace-id-value"
def test_get_chain_id_from_headers_generic_vendor_session_id():
"""get_chain_id_from_headers picks up any x-<vendor>-session-id with a valid value."""
from litellm.proxy.litellm_pre_call_utils import get_chain_id_from_headers
assert (
get_chain_id_from_headers(
{"x-claude-code-session-id": "e96634a3-fa28-4083-b354-55542e2dca01"}
)
== "e96634a3-fa28-4083-b354-55542e2dca01"
)
# Short / non-alphanumeric values should be ignored
assert get_chain_id_from_headers({"x-foo-session-id": "short"}) is None
assert get_chain_id_from_headers({"x-foo-session-id": "has spaces!!"}) is None
# Explicit headers still take precedence
assert (
get_chain_id_from_headers(
{
"x-litellm-trace-id": "explicit-id-value",
"x-claude-code-session-id": "e96634a3-fa28-4083-b354-55542e2dca01",
}
)
== "explicit-id-value"
)
def test_get_internal_user_header_from_mapping_returns_expected_header():
mappings = [
{"header_name": "X-OpenWebUI-User-Id", "litellm_user_role": "internal_user"},
+317 -8
View File
@@ -3,6 +3,7 @@ Unit tests for litellm.compress().
"""
import os
import importlib
import pytest
@@ -12,6 +13,10 @@ from litellm.compression.scoring.embedding_scorer import embedding_score_message
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
from litellm.types.utils import CallTypes
CALL_TYPE = CallTypes.completion
ANTHROPIC_CALL_TYPE = CallTypes.anthropic_messages
# ---------------------------------------------------------------------------
@@ -149,7 +154,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", call_type=CALL_TYPE)
assert result["messages"] == messages
assert result["cache"] == {}
assert result["tools"] == []
@@ -178,6 +183,7 @@ def test_compress_above_trigger():
result = litellm.compress(
big_messages,
model="gpt-4o",
call_type=CALL_TYPE,
compression_trigger=1000,
compression_target=500,
)
@@ -189,13 +195,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",
call_type=ANTHROPIC_CALL_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", call_type=CALL_TYPE, compression_trigger=1000
)
assert result["messages"][0]["role"] == "system"
assert "System prompt" in result["messages"][0]["content"]
@@ -205,7 +260,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", call_type=CALL_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 +273,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", call_type=CALL_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 +288,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", call_type=CALL_TYPE, compression_trigger=1000
)
if result["tools"]:
tool_desc = result["tools"][0]["function"]["description"]
for key in result["cache"]:
@@ -242,11 +303,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", call_type=CALL_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",
call_type=ANTHROPIC_CALL_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_call_type_is_completion():
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 +394,7 @@ def test_compress_forwards_embedding_model_params(monkeypatch):
{"role": "user", "content": "Fix auth"},
],
model="gpt-4o",
call_type=CALL_TYPE,
compression_trigger=1000,
embedding_model="text-embedding-3-small",
embedding_model_params={"api_base": "https://example-embeddings.test"},
@@ -326,6 +452,7 @@ def test_embedding_scorer():
{"role": "user", "content": "Fix auth"},
],
model="gpt-4o",
call_type=CALL_TYPE,
compression_trigger=1000,
embedding_model="text-embedding-3-small",
)
@@ -346,8 +473,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", call_type=CALL_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"]
@@ -356,3 +484,184 @@ def test_simple_compression(final_user_message, expected_content):
assert "Unrelated cooking recipes " not in result["messages"][1]["content"]
else:
raise ValueError(f"Unexpected expected_content: {expected_content}")
def test_compress_anthropic_drops_irrelevant_tool_exchange_span(monkeypatch):
compress_module = importlib.import_module("litellm.compression.compress")
def fake_bm25_score_messages(query, messages):
assert "final query" in query
assert len(messages) == 5
# Prefer idx=0 and de-prioritize the tool exchange span (idx=1,2)
return [0.95, 0.01, 0.02, 0.8, 1.0]
def fake_token_counter(model, messages=None, text=None):
if messages is not None:
return 1000
if text is None:
return 0
if "final query" in text:
return 50
if "assistant_tail" in text:
return 20
if "other_blob" in text:
return 220
if "tool_payload_relevant" in text:
return 200
if text == "":
return 1
return 10
monkeypatch.setattr(
compress_module, "bm25_score_messages", fake_bm25_score_messages
)
monkeypatch.setattr(compress_module, "token_counter", fake_token_counter)
messages = [
{"role": "user", "content": "other_blob " * 300},
{
"role": "assistant",
"content": [
{
"type": "tool_use",
"id": "toolu_drop",
"name": "litellm_content_retrieve",
"input": {"key": "message_1"},
}
],
},
{
"role": "user",
"content": [
{
"type": "tool_result",
"tool_use_id": "toolu_drop",
"content": [{"type": "text", "text": "tool_payload_relevant"}],
}
],
},
{"role": "assistant", "content": "assistant_tail"},
{"role": "user", "content": "final query"},
]
result = litellm.compress(
messages=messages,
model="claude-sonnet-4-20250514",
call_type=ANTHROPIC_CALL_TYPE,
compression_trigger=100,
compression_target=280,
)
# idx=1,2 should be dropped atomically (no orphan tool blocks left behind)
assert len(result["messages"]) == 3
assert result["messages"][0]["role"] == "user"
assert "other_blob" in result["messages"][0]["content"]
assert result["messages"][1]["content"] == "assistant_tail"
assert result["messages"][2]["content"] == "final query"
assert result["cache"] == {}
def test_compress_anthropic_keeps_relevant_tool_exchange_span(monkeypatch):
compress_module = importlib.import_module("litellm.compression.compress")
def fake_bm25_score_messages(query, messages):
assert "final query" in query
assert len(messages) == 5
# Prefer the tool exchange span over idx=0
return [0.05, 0.01, 0.92, 0.8, 1.0]
def fake_token_counter(model, messages=None, text=None):
if messages is not None:
return 1000
if text is None:
return 0
if "final query" in text:
return 50
if "assistant_tail" in text:
return 20
if "other_blob" in text:
return 220
if "tool_payload_relevant" in text:
return 200
if text == "":
return 1
return 10
monkeypatch.setattr(
compress_module, "bm25_score_messages", fake_bm25_score_messages
)
monkeypatch.setattr(compress_module, "token_counter", fake_token_counter)
messages = [
{"role": "user", "content": "other_blob " * 300},
{
"role": "assistant",
"content": [
{
"type": "tool_use",
"id": "toolu_keep",
"name": "litellm_content_retrieve",
"input": {"key": "message_1"},
}
],
},
{
"role": "user",
"content": [
{
"type": "tool_result",
"tool_use_id": "toolu_keep",
"content": [{"type": "text", "text": "tool_payload_relevant"}],
}
],
},
{"role": "assistant", "content": "assistant_tail"},
{"role": "user", "content": "final query"},
]
result = litellm.compress(
messages=messages,
model="claude-sonnet-4-20250514",
call_type=ANTHROPIC_CALL_TYPE,
compression_trigger=100,
compression_target=280,
)
assert len(result["messages"]) == 5
assert result["messages"][1]["role"] == "assistant"
assert result["messages"][2]["role"] == "user"
# idx=0 should be compressed instead
assert "litellm_content_retrieve" in result["messages"][0]["content"]
assert len(result["cache"]) == 1
def test_compress_anthropic_malformed_tool_sequence_passes_through():
messages = [
{"role": "user", "content": "other_blob " * 300},
{
"role": "assistant",
"content": [
{
"type": "tool_use",
"id": "toolu_broken",
"name": "litellm_content_retrieve",
"input": {"key": "message_1"},
}
],
},
{"role": "user", "content": [{"type": "text", "text": "missing tool_result"}]},
{"role": "user", "content": "final query"},
]
result = litellm.compress(
messages=messages,
model="claude-sonnet-4-20250514",
call_type=ANTHROPIC_CALL_TYPE,
compression_trigger=100,
compression_target=280,
)
assert result["messages"] == messages
assert result["cache"] == {}
assert result["tools"] == []
assert result["compression_skipped_reason"] == "invalid_anthropic_tool_sequence"