mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-01 22:22:22 +00:00
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:
co-authored by
Claude Opus 4.6
parent
a19bff4ca6
commit
386f334fee
@@ -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`
|
||||
@@ -536,6 +536,7 @@ const sidebars = {
|
||||
description: "Modify requests, responses, and more",
|
||||
items: [
|
||||
"proxy/call_hooks",
|
||||
"proxy/agentic_loop_hook",
|
||||
"proxy/rules",
|
||||
]
|
||||
},
|
||||
|
||||
@@ -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
@@ -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
|
||||
@@ -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"""
|
||||
|
||||
+320
@@ -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
|
||||
@@ -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)}"
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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]]
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
+364
@@ -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]"
|
||||
)
|
||||
+56
-1
@@ -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.
|
||||
|
||||
+792
@@ -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"},
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user