mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge pull request #34290 from eugene-yao-zocdoc/litellm_fix_anthropic_midturn_interjections
fix(anthropic): preserve midturn system corrections
This commit is contained in:
commit
1a8cd8a078
9 changed files with 1437 additions and 54 deletions
|
|
@ -14,6 +14,7 @@ Pattern Overview:
|
|||
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from copy import deepcopy
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
|
||||
|
|
@ -110,14 +111,10 @@ EMPTY_EXTRACTED_INPUT: Final = ExtractedInput(scanned=(), images=())
|
|||
|
||||
|
||||
class AnthropicMessagesHandler(BaseTranslation):
|
||||
"""
|
||||
Handler for processing Anthropic messages with guardrails.
|
||||
"""Process Anthropic messages with guardrails.
|
||||
|
||||
This class provides methods to:
|
||||
1. Process input messages (pre-call hook)
|
||||
2. Process output responses (post-call hook)
|
||||
|
||||
Methods can be overridden to customize behavior for different message formats.
|
||||
In-sequence system entries are untrusted client input. This handler scans and preserves
|
||||
them through guardrail rewrites; downstream provider handling is out of scope.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
|
|
@ -331,16 +328,30 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
skip_tool: Final = effective_skip_tool_message_for_guardrail(guardrail_to_apply)
|
||||
scan_only_tool_results: Final = effective_scan_only_tool_results_for_guardrail(guardrail_to_apply)
|
||||
|
||||
chat_completion_compatible_request: Final = self._translate_to_openai(data)
|
||||
# Exclude only the trusted top-level prompt. In-sequence system entries are untrusted
|
||||
# and must stay aligned with texts_to_check for positional masking. When the top-level
|
||||
# prompt is included, the pre-existing count mismatch disables positional masking.
|
||||
translation_source: Final = { # mutable-ok: API message payload
|
||||
key: value for key, value in data.items() if key != "system"
|
||||
}
|
||||
chat_completion_compatible_request: Final = self._translate_to_openai(translation_source)
|
||||
|
||||
full_structured_messages: Final = cast(
|
||||
list[AllMessageValues],
|
||||
chat_completion_compatible_request.get("messages", []),
|
||||
)
|
||||
has_midturn_system_message: Final = any(
|
||||
str(message.get("role") or "").lower() == "system" for message in full_structured_messages
|
||||
)
|
||||
hoisted_system_message: Final = None if skip_system else self._hoisted_top_level_system_message(data)
|
||||
if hoisted_system_message is not None:
|
||||
full_structured_messages.insert(0, hoisted_system_message)
|
||||
# skip_system already excluded the trusted top-level prompt (it is simply not hoisted);
|
||||
# in-sequence system entries are untrusted and always stay in scope.
|
||||
scoped_message_indices: Final = scoped_structured_message_indices(
|
||||
full_structured_messages,
|
||||
scan_only_tool_results=scan_only_tool_results,
|
||||
skip_system=skip_system,
|
||||
skip_system=False,
|
||||
skip_tool=skip_tool,
|
||||
)
|
||||
structured_messages: Final = [full_structured_messages[index] for index in scoped_message_indices]
|
||||
|
|
@ -422,6 +433,8 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
scoped_indices=scoped_message_indices,
|
||||
guardrailed_scoped=guardrailed_structured_messages,
|
||||
),
|
||||
hoisted_system_message=hoisted_system_message,
|
||||
preserve_system_messages=has_midturn_system_message,
|
||||
)
|
||||
else:
|
||||
# Step 3: Map guardrail responses back to original message structure
|
||||
|
|
@ -435,36 +448,150 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
|
||||
return data
|
||||
|
||||
@staticmethod
|
||||
def _write_back_structured_messages(data: dict, structured_messages: list) -> None:
|
||||
"""Convert compressed structured_messages back to Anthropic format and write to data.
|
||||
def _hoisted_top_level_system_message(
|
||||
self, data: dict
|
||||
) -> AllMessageValues | None: # mutable-ok: API message payload
|
||||
"""Return the system message produced by translating the top-level prompt."""
|
||||
system: Final = data.get("system")
|
||||
if not system:
|
||||
return None
|
||||
probe: Final = self._translate_to_openai(
|
||||
{ # mutable-ok: API message payload
|
||||
"model": data.get("model") or "",
|
||||
"messages": [], # mutable-ok: API message payload
|
||||
"system": system,
|
||||
}
|
||||
)
|
||||
hoisted: Final = probe.get("messages") or [] # mutable-ok: API message payload
|
||||
return hoisted[0] if hoisted else None
|
||||
|
||||
``anthropic_messages_pt`` merges every run of consecutive user/tool rows
|
||||
into a single message, so a turn carrying only tool results and the user
|
||||
turn that follows it come back fused, and the request the model sees no
|
||||
longer has the boundaries the client sent. Converting a row at a time
|
||||
would keep them apart but breaks tool pairing: an assistant row whose
|
||||
tool results sit outside its own call reads as an orphaned tool call,
|
||||
and under ``modify_params`` the sanitizer answers it with a synthetic
|
||||
"tool execution skipped" result and drops the real one. Converting each
|
||||
assistant row together with the tool rows that answer it, and every
|
||||
other row on its own, satisfies both.
|
||||
"""
|
||||
@staticmethod
|
||||
def _openai_system_message_to_anthropic(
|
||||
message: dict[str, Any],
|
||||
) -> dict[str, Any] | None: # mutable-ok: API message payload
|
||||
"""Convert an OpenAI system message to the client's Anthropic-shaped entry."""
|
||||
content: Final = message.get("content")
|
||||
if isinstance(content, str):
|
||||
return (
|
||||
{"role": "system", "content": content} if content else None # mutable-ok: API message payload
|
||||
) # mutable-ok: API message payload
|
||||
if not isinstance(content, list):
|
||||
return None
|
||||
blocks: Final[list[dict[str, Any]]] = [] # mutable-ok: API message payload
|
||||
for block in content:
|
||||
if not isinstance(block, dict) or block.get("type") != "text":
|
||||
continue
|
||||
text = block.get("text")
|
||||
if not isinstance(text, str) or not text:
|
||||
continue
|
||||
anthropic_block: dict[str, Any] = { # mutable-ok: API message payload
|
||||
"type": "text",
|
||||
"text": text,
|
||||
} # mutable-ok: API message payload
|
||||
cache_control = block.get("cache_control")
|
||||
if cache_control:
|
||||
anthropic_block["cache_control"] = deepcopy(cache_control)
|
||||
blocks.append(anthropic_block)
|
||||
return (
|
||||
{"role": "system", "content": blocks} if blocks else None # mutable-ok: API message payload
|
||||
) # mutable-ok: API message payload
|
||||
|
||||
@staticmethod
|
||||
def _is_hoisted_top_level_system(message: object, hoisted_system_message: object) -> bool:
|
||||
"""Match the hoisted prompt by identity, or by value after serialization."""
|
||||
if hoisted_system_message is None:
|
||||
return False
|
||||
if message is hoisted_system_message:
|
||||
return True
|
||||
return (
|
||||
isinstance(message, dict) and isinstance(hoisted_system_message, dict) and message == hoisted_system_message
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _is_system(message: object) -> bool:
|
||||
"""Whether the row is an in-sequence system message."""
|
||||
return isinstance(message, dict) and str(message.get("role") or "").lower() == "system"
|
||||
|
||||
@staticmethod
|
||||
def _defer_systems_inside_tool_exchanges(
|
||||
structured_messages: list, # mutable-ok: API message payload
|
||||
) -> list:
|
||||
"""Hold a system row until the tool exchange around it completes so the call/result pair converts together."""
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import group_tool_exchanges
|
||||
|
||||
non_system_positions: Final[list[int]] = [
|
||||
index
|
||||
for index, message in enumerate(structured_messages)
|
||||
if not AnthropicMessagesHandler._is_system(message)
|
||||
]
|
||||
exchange_end_for_start: Final[dict[int, int]] = {
|
||||
non_system_positions[group[0]]: non_system_positions[group[-1]]
|
||||
for group in group_tool_exchanges([structured_messages[index] for index in non_system_positions])
|
||||
if len(group) > 1
|
||||
}
|
||||
ordered: Final[list] = [] # mutable-ok: API message payload
|
||||
deferred_systems: Final[list] = [] # mutable-ok: API message payload
|
||||
open_exchange_end = -1 # rebind-ok: advances to the enclosing exchange's last index
|
||||
for index, message in enumerate(structured_messages):
|
||||
if AnthropicMessagesHandler._is_system(message) and index < open_exchange_end:
|
||||
deferred_systems.append(message)
|
||||
continue
|
||||
open_exchange_end = exchange_end_for_start.get(index, open_exchange_end)
|
||||
ordered.append(message)
|
||||
if index >= open_exchange_end and deferred_systems:
|
||||
ordered.extend(deferred_systems)
|
||||
deferred_systems.clear()
|
||||
ordered.extend(deferred_systems)
|
||||
return ordered
|
||||
|
||||
@staticmethod
|
||||
def _write_back_structured_messages(
|
||||
data: dict, # mutable-ok: API message payload
|
||||
structured_messages: list, # mutable-ok: API message payload
|
||||
hoisted_system_message: object = None,
|
||||
preserve_system_messages: bool = False,
|
||||
) -> None:
|
||||
"""Write a guardrail's structured-message rewrite back without losing corrections."""
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import (
|
||||
anthropic_messages_pt,
|
||||
group_tool_exchanges,
|
||||
)
|
||||
|
||||
_is_system: Final = AnthropicMessagesHandler._is_system
|
||||
model: Final = str(data.get("model") or "")
|
||||
non_system: Final = [m for m in structured_messages if m.get("role") != "system"]
|
||||
groups: Final = tuple([non_system[index] for index in group] for group in group_tool_exchanges(non_system)) or (
|
||||
non_system,
|
||||
)
|
||||
converted: Final = [
|
||||
message
|
||||
for group in groups
|
||||
for message in anthropic_messages_pt(messages=group, model=model, llm_provider="anthropic")
|
||||
]
|
||||
converted: Final[list] = [] # mutable-ok: API message payload
|
||||
|
||||
def _convert_run(run: list) -> None: # mutable-ok: API message payload
|
||||
for group in group_tool_exchanges(run):
|
||||
converted.extend(
|
||||
anthropic_messages_pt(
|
||||
messages=[run[index] for index in group], # mutable-ok: API message payload
|
||||
model=model,
|
||||
llm_provider="anthropic",
|
||||
)
|
||||
)
|
||||
|
||||
ordered: Final = AnthropicMessagesHandler._defer_systems_inside_tool_exchanges(structured_messages)
|
||||
run: Final[list] = [] # mutable-ok: API message payload
|
||||
hoisted_dropped = False # rebind-ok: flips once the hoisted prompt is dropped
|
||||
for message in ordered:
|
||||
if not _is_system(message):
|
||||
run.append(message)
|
||||
continue
|
||||
_convert_run(run)
|
||||
run.clear()
|
||||
if not hoisted_dropped and AnthropicMessagesHandler._is_hoisted_top_level_system(
|
||||
message, hoisted_system_message
|
||||
):
|
||||
hoisted_dropped = True
|
||||
continue
|
||||
if preserve_system_messages:
|
||||
anthropic_system = AnthropicMessagesHandler._openai_system_message_to_anthropic(message)
|
||||
if anthropic_system is not None:
|
||||
converted.append(anthropic_system)
|
||||
_convert_run(run)
|
||||
if not any(not _is_system(message) for message in converted):
|
||||
converted.extend(anthropic_messages_pt(messages=[], model=model, llm_provider="anthropic"))
|
||||
for msg in converted:
|
||||
content = msg.get("content")
|
||||
if isinstance(content, list):
|
||||
|
|
@ -473,6 +600,31 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
block.pop("cache_control", None)
|
||||
data["messages"] = converted
|
||||
|
||||
@staticmethod
|
||||
def _extract_midturn_system_text(
|
||||
message: dict[str, Any], # mutable-ok: API message payload
|
||||
msg_idx: int,
|
||||
) -> ExtractedInput:
|
||||
"""Match the adapter's filtering so positional guardrail write-back stays aligned."""
|
||||
content: Final = message.get("content")
|
||||
if isinstance(content, str):
|
||||
if not content:
|
||||
return EMPTY_EXTRACTED_INPUT
|
||||
return ExtractedInput(scanned=(ScannedText(content, MessageContentTarget(msg_idx)),), images=())
|
||||
if not isinstance(content, list):
|
||||
return EMPTY_EXTRACTED_INPUT
|
||||
return ExtractedInput(
|
||||
scanned=tuple(
|
||||
ScannedText(text_str, ContentBlockTextTarget(msg_idx, content_idx))
|
||||
for content_idx, content_item in enumerate(content)
|
||||
if isinstance(content_item, dict)
|
||||
and content_item.get("type") == "text"
|
||||
and isinstance(text_str := content_item.get("text"), str)
|
||||
and text_str
|
||||
),
|
||||
images=(),
|
||||
)
|
||||
|
||||
def extract_request_tool_names(self, data: dict) -> list[str]:
|
||||
"""Extract tool names from Anthropic messages request (tools[].name)."""
|
||||
names: Final[list[str]] = []
|
||||
|
|
@ -490,11 +642,17 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
skip_tool_message: bool = False,
|
||||
scan_only_tool_results: bool = False,
|
||||
) -> ExtractedInput:
|
||||
"""Extract text content and images from a message.
|
||||
|
||||
In-sequence system entries are scanned even when ``skip_system_message`` is set:
|
||||
that flag covers only the trusted top-level prompt, which never appears here.
|
||||
"""
|
||||
Extract text content and images from a message.
|
||||
"""
|
||||
role: Final = str(message.get("role") or "").lower()
|
||||
if (skip_system_message and role == "system") or (skip_tool_message and role == "tool"):
|
||||
role: Final = str(message.get("role") or "")
|
||||
if role == "system":
|
||||
if scan_only_tool_results:
|
||||
return EMPTY_EXTRACTED_INPUT
|
||||
return cls._extract_midturn_system_text(message=message, msg_idx=msg_idx)
|
||||
if skip_tool_message and role.lower() == "tool":
|
||||
return EMPTY_EXTRACTED_INPUT
|
||||
|
||||
content: Final = message.get("content", None)
|
||||
|
|
|
|||
|
|
@ -74,12 +74,12 @@ from litellm.llms.anthropic.experimental_pass_through.context_management import
|
|||
)
|
||||
from litellm.types.llms.anthropic import (
|
||||
ANTHROPIC_HOSTED_TOOLS,
|
||||
AllAnthropicPassThroughMessageValues,
|
||||
AllAnthropicToolsValues,
|
||||
AnthopicMessagesAssistantMessageParam,
|
||||
AnthropicFinishReason,
|
||||
AnthropicMessagesRequest,
|
||||
AnthropicMessagesSystemMessageParam,
|
||||
AnthropicMessagesToolChoice,
|
||||
AnthropicMessagesUserMessageParam,
|
||||
AnthropicResponseContentBlockRedactedThinking,
|
||||
AnthropicResponseContentBlockText,
|
||||
AnthropicResponseContentBlockThinking,
|
||||
|
|
@ -343,7 +343,7 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
|
||||
def translate_anthropic_messages_to_openai(
|
||||
self,
|
||||
messages: list[AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam],
|
||||
messages: list[AllAnthropicPassThroughMessageValues],
|
||||
model: str | None = None,
|
||||
) -> list:
|
||||
new_messages: Final[list[AllMessageValues]] = []
|
||||
|
|
@ -351,6 +351,11 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
user_message: ChatCompletionUserMessage | None = None
|
||||
tool_message_list: list[ChatCompletionToolMessage] = []
|
||||
new_user_content_list: list[ChatCompletionTextObject | ChatCompletionImageObject] = []
|
||||
if m["role"] == "system":
|
||||
system_message = self._translate_midturn_system_message_to_openai(m, model)
|
||||
if system_message is not None:
|
||||
new_messages.append(system_message)
|
||||
continue
|
||||
## USER MESSAGE ##
|
||||
if m["role"] == "user":
|
||||
## translate user message
|
||||
|
|
@ -848,6 +853,29 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
for def_schema in schema[key].values():
|
||||
LiteLLMAnthropicMessagesAdapter._add_additional_properties_false(def_schema)
|
||||
|
||||
def _translate_midturn_system_message_to_openai(
|
||||
self,
|
||||
message: AnthropicMessagesSystemMessageParam,
|
||||
model: str | None,
|
||||
) -> ChatCompletionSystemMessage | None:
|
||||
"""Translate an in-sequence system entry without changing its role or position."""
|
||||
content: Final = message.get("content")
|
||||
if isinstance(content, str):
|
||||
return ChatCompletionSystemMessage(role="system", content=content) if content else None
|
||||
if not isinstance(content, list):
|
||||
return None
|
||||
text_parts: Final[list[ChatCompletionTextObject]] = [] # mutable-ok: API message payload
|
||||
for block in content:
|
||||
if not isinstance(block, dict) or block.get("type") != "text": # pyright: ignore[reportUnnecessaryIsInstance] # untrusted client payload
|
||||
continue
|
||||
text = block.get("text")
|
||||
if not text:
|
||||
continue
|
||||
text_obj = ChatCompletionTextObject(type="text", text=text)
|
||||
self._add_cache_control_if_applicable(block, text_obj, model)
|
||||
text_parts.append(text_obj)
|
||||
return ChatCompletionSystemMessage(role="system", content=text_parts) if text_parts else None
|
||||
|
||||
def _add_system_message_to_messages(
|
||||
self,
|
||||
new_messages: list[AllMessageValues],
|
||||
|
|
@ -1049,8 +1077,8 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
tool_name_mapping: dict[str, str] = {}
|
||||
|
||||
## CONVERT ANTHROPIC MESSAGES TO OPENAI
|
||||
messages_list: Final[list[AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam]] = cast(
|
||||
list[AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam],
|
||||
messages_list: Final[list[AllAnthropicPassThroughMessageValues]] = cast(
|
||||
list[AllAnthropicPassThroughMessageValues],
|
||||
anthropic_message_request["messages"],
|
||||
)
|
||||
new_messages = self.translate_anthropic_messages_to_openai(
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ path used for OpenAI and Azure models.
|
|||
"""
|
||||
|
||||
import json
|
||||
from collections.abc import Iterable
|
||||
from typing import Any, Final, cast
|
||||
|
||||
from litellm.litellm_core_utils.reasoning_effort_utils import (
|
||||
|
|
@ -15,15 +16,15 @@ from litellm.llms.anthropic.experimental_pass_through.utils import (
|
|||
is_reasoning_auto_summary_enabled,
|
||||
)
|
||||
from litellm.types.llms.anthropic import (
|
||||
AllAnthropicPassThroughMessageValues,
|
||||
AllAnthropicToolsValues,
|
||||
AnthopicMessagesAssistantMessageParam,
|
||||
AnthropicFinishReason,
|
||||
AnthropicMessagesRequest,
|
||||
AnthropicMessagesToolChoice,
|
||||
AnthropicMessagesUserMessageParam,
|
||||
AnthropicResponseContentBlockText,
|
||||
AnthropicResponseContentBlockThinking,
|
||||
AnthropicResponseContentBlockToolUse,
|
||||
AnthropicSystemMessageContent,
|
||||
)
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import (
|
||||
AnthropicMessagesResponse,
|
||||
|
|
@ -72,14 +73,32 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
|
|||
return source.get("url")
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _translate_midturn_system_content_to_responses(
|
||||
content: str | Iterable[AnthropicSystemMessageContent],
|
||||
) -> list[dict[str, str]]: # mutable-ok: API message payload
|
||||
"""Convert in-sequence system content to Responses input-text parts."""
|
||||
if isinstance(content, str):
|
||||
return (
|
||||
[{"type": "input_text", "text": content}] if content else [] # mutable-ok: API message payload
|
||||
) # mutable-ok: API message payload
|
||||
if not isinstance(content, list):
|
||||
return [] # mutable-ok: API message payload
|
||||
return [ # mutable-ok: API message payload
|
||||
{"type": "input_text", "text": text} # mutable-ok: API message payload
|
||||
for block in content
|
||||
if isinstance(block, dict) and block.get("type") == "text" and (text := block.get("text")) # pyright: ignore[reportUnnecessaryIsInstance] # untrusted client payload
|
||||
]
|
||||
|
||||
def translate_messages_to_responses_input(
|
||||
self,
|
||||
messages: list[AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam],
|
||||
messages: list[AllAnthropicPassThroughMessageValues],
|
||||
) -> list[dict[str, Any]]:
|
||||
"""
|
||||
Convert Anthropic messages list to Responses API `input` items.
|
||||
|
||||
Mapping:
|
||||
system text -> message(role=system, input_text)
|
||||
user text -> message(role=user, input_text)
|
||||
user image -> message(role=user, input_image)
|
||||
user tool_result -> function_call_output
|
||||
|
|
@ -89,6 +108,18 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
|
|||
input_items: Final[list[dict[str, Any]]] = []
|
||||
|
||||
for m in messages:
|
||||
if m["role"] == "system":
|
||||
system_parts = self._translate_midturn_system_content_to_responses(m.get("content"))
|
||||
if system_parts:
|
||||
input_items.append(
|
||||
{ # mutable-ok: API message payload
|
||||
"type": "message",
|
||||
"role": "system",
|
||||
"content": system_parts,
|
||||
}
|
||||
)
|
||||
continue
|
||||
|
||||
role = m["role"]
|
||||
content = m.get("content")
|
||||
|
||||
|
|
@ -300,7 +331,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
|
|||
"""
|
||||
model: Final[str] = anthropic_request["model"]
|
||||
messages_list: Final = cast(
|
||||
list[AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam],
|
||||
list[AllAnthropicPassThroughMessageValues],
|
||||
anthropic_request["messages"],
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -749,7 +749,10 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up
|
|||
"When True, unified guardrails skip system-role messages when building "
|
||||
"evaluation inputs (texts and structured_messages). When False, system "
|
||||
"messages are included even if litellm_settings sets a global skip. When "
|
||||
"None, use the global litellm.skip_system_message_in_guardrail setting."
|
||||
"None, use the global litellm.skip_system_message_in_guardrail setting. "
|
||||
"For Anthropic /v1/messages, the flag applies only to the trusted top-level "
|
||||
"system prompt. In-sequence system entries are untrusted client input and remain "
|
||||
"in texts and structured_messages."
|
||||
),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
from collections.abc import Iterable
|
||||
from enum import Enum
|
||||
from typing import Any, Final, Literal
|
||||
from typing import Any, Final, Literal, TypeAlias
|
||||
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from typing_extensions import NotRequired, Required, TypedDict
|
||||
|
|
@ -348,8 +348,18 @@ class AnthropicSystemMessageContent(TypedDict, total=False):
|
|||
cache_control: dict | ChatCompletionCachedContent | None
|
||||
|
||||
|
||||
class AnthropicMessagesSystemMessageParam(TypedDict, total=False):
|
||||
role: Required[Literal["system"]]
|
||||
content: Required[str | Iterable[AnthropicSystemMessageContent]]
|
||||
|
||||
|
||||
AllAnthropicMessageValues = AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam
|
||||
|
||||
# System is not a native Anthropic message role; only pass-through adapters use this union.
|
||||
AllAnthropicPassThroughMessageValues: TypeAlias = (
|
||||
AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam | AnthropicMessagesSystemMessageParam
|
||||
)
|
||||
|
||||
|
||||
class AnthropicMessagesRequestOptionalParams(TypedDict, total=False):
|
||||
max_tokens: int | None
|
||||
|
|
|
|||
|
|
@ -76,6 +76,51 @@ class MockRecordingGuardrail(CustomGuardrail):
|
|||
return inputs
|
||||
|
||||
|
||||
class MockMaskingGuardrail(CustomGuardrail):
|
||||
"""Capture request inputs and mask one known prohibited value."""
|
||||
|
||||
def __init__(self, skip_system_message_in_guardrail: Optional[bool] = True):
|
||||
super().__init__(guardrail_name="masking-test")
|
||||
self.skip_system_message_in_guardrail = skip_system_message_in_guardrail
|
||||
self.inputs: Optional[GenericGuardrailAPIInputs] = None
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional[Any] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
self.inputs = inputs.copy()
|
||||
masked_inputs = inputs.copy()
|
||||
masked_inputs["texts"] = [
|
||||
"[MASKED]" if text == "prohibited correction" else text for text in inputs.get("texts", [])
|
||||
]
|
||||
return masked_inputs
|
||||
|
||||
|
||||
class MockCompactingGuardrail(CustomGuardrail):
|
||||
"""Stand in for a compaction guardrail that rewrites `structured_messages` wholesale."""
|
||||
|
||||
def __init__(self, replacement_messages: list):
|
||||
super().__init__(guardrail_name="compacting-test")
|
||||
self.replacement_messages = replacement_messages
|
||||
self.inputs: Optional[GenericGuardrailAPIInputs] = None
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional[Any] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
self.inputs = inputs.copy()
|
||||
rewritten = inputs.copy()
|
||||
# A new list object -- this is what signals a rewrite to the handler.
|
||||
rewritten["structured_messages"] = list(self.replacement_messages)
|
||||
return rewritten
|
||||
|
||||
|
||||
class TestAnthropicMessagesHandlerStreamingRequestData:
|
||||
"""Post-call guardrails on streaming /v1/messages receive the response and identity metadata"""
|
||||
|
||||
|
|
@ -211,6 +256,704 @@ class TestAnthropicMessagesHandlerInputProcessing:
|
|||
assert data.get("litellm_metadata", {}).get("guardrails")
|
||||
assert guardrail.dynamic_params == {"policy_id": "policy-123"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_midturn_system_correction_is_guardrailed_when_top_level_system_is_skipped(
|
||||
self,
|
||||
):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockMaskingGuardrail()
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"system": "trusted top-level system prompt",
|
||||
"messages": [
|
||||
{"role": "user", "content": "safe text"},
|
||||
{
|
||||
"role": "system",
|
||||
"content": [
|
||||
{"type": "unsupported", "text": "discarded text"},
|
||||
{"type": "text", "text": "prohibited correction"},
|
||||
],
|
||||
},
|
||||
],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert guardrail.inputs is not None
|
||||
assert guardrail.inputs["texts"] == ["safe text", "prohibited correction"]
|
||||
assert "trusted top-level system prompt" not in guardrail.inputs["texts"]
|
||||
assert data["messages"][1]["content"][0]["text"] == "discarded text"
|
||||
assert data["messages"][1]["content"][1]["text"] == "[MASKED]"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_string_midturn_system_correction_is_guardrailed(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockMaskingGuardrail()
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"messages": [{"role": "system", "content": "prohibited correction"}],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert guardrail.inputs is not None
|
||||
assert guardrail.inputs["texts"] == ["prohibited correction"]
|
||||
assert data["messages"][0]["content"] == "[MASKED]"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unsupported_midturn_system_content_is_not_guardrailed(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockMaskingGuardrail()
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"messages": [
|
||||
{
|
||||
"role": "system",
|
||||
"content": [{"type": "image", "source": {"type": "url"}}],
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert guardrail.inputs is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_skip_system_message_excludes_only_hoisted_top_level_system(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockMaskingGuardrail()
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"system": "trusted top-level system prompt",
|
||||
"messages": [
|
||||
{"role": "user", "content": "safe text"},
|
||||
{"role": "system", "content": "prohibited correction"},
|
||||
{"role": "user", "content": "continue"},
|
||||
],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert guardrail.inputs is not None
|
||||
structured = guardrail.inputs["structured_messages"]
|
||||
assert [m["role"] for m in structured] == ["user", "system", "user"]
|
||||
assert structured[1]["content"] == "prohibited correction"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_default_skip_false_scans_midturn_system_and_hoists_top_level_system(
|
||||
self,
|
||||
):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockMaskingGuardrail(skip_system_message_in_guardrail=None)
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"system": "trusted top-level system prompt",
|
||||
"messages": [
|
||||
{"role": "user", "content": "safe text"},
|
||||
{"role": "system", "content": "prohibited correction"},
|
||||
],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert guardrail.inputs is not None
|
||||
assert guardrail.inputs["texts"] == ["safe text", "prohibited correction"]
|
||||
structured = guardrail.inputs["structured_messages"]
|
||||
assert [m["role"] for m in structured] == ["system", "user", "system"]
|
||||
assert structured[0]["content"] == "trusted top-level system prompt"
|
||||
assert data["messages"][1]["content"] == "[MASKED]"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_masking_slice_is_unavailable_when_top_level_system_is_included(
|
||||
self,
|
||||
):
|
||||
from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
|
||||
BedrockGuardrail,
|
||||
)
|
||||
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockMaskingGuardrail(skip_system_message_in_guardrail=None)
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"system": "trusted top-level system prompt",
|
||||
"messages": [
|
||||
{"role": "user", "content": "safe text"},
|
||||
{"role": "system", "content": "prohibited correction"},
|
||||
{"role": "user", "content": "latest question"},
|
||||
],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert guardrail.inputs is not None
|
||||
texts = guardrail.inputs["texts"]
|
||||
structured = guardrail.inputs["structured_messages"]
|
||||
|
||||
bedrock = BedrockGuardrail(guardrailIdentifier="gi", guardrailVersion="1")
|
||||
assert sum(bedrock._count_message_texts(m) for m in structured) == len(texts) + 1
|
||||
latest_user_index = bedrock._find_latest_message_index(structured, target_role="user")
|
||||
assert (
|
||||
bedrock._locate_message_texts_slice(
|
||||
structured_messages=structured,
|
||||
target_index=latest_user_index,
|
||||
texts=texts,
|
||||
)
|
||||
is None
|
||||
)
|
||||
assert (
|
||||
bedrock._merge_masked_texts(
|
||||
masked_texts=["{MASKED}"],
|
||||
texts=texts,
|
||||
scanned_slice=None,
|
||||
scanned_role_subset=True,
|
||||
)
|
||||
== texts
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("skip_system_message_in_guardrail", [True, None])
|
||||
async def test_midturn_system_text_extraction_matches_translation_in_both_skip_modes(
|
||||
self,
|
||||
skip_system_message_in_guardrail: Optional[bool],
|
||||
):
|
||||
from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
|
||||
BedrockGuardrail,
|
||||
)
|
||||
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockMaskingGuardrail(skip_system_message_in_guardrail=skip_system_message_in_guardrail)
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"messages": [
|
||||
{"role": "user", "content": "safe text"},
|
||||
{
|
||||
"role": "system",
|
||||
"content": [
|
||||
{"type": "text", "text": ""},
|
||||
{"type": "image", "source": {"type": "url", "url": "https://example.com/a.png"}},
|
||||
{"type": "text", "text": "prohibited correction"},
|
||||
],
|
||||
},
|
||||
],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert guardrail.inputs is not None
|
||||
texts = guardrail.inputs["texts"]
|
||||
structured = guardrail.inputs["structured_messages"]
|
||||
assert texts == ["safe text", "prohibited correction"]
|
||||
bedrock = BedrockGuardrail(guardrailIdentifier="gi", guardrailVersion="1")
|
||||
assert sum(bedrock._count_message_texts(m) for m in structured) == len(texts)
|
||||
assert data["messages"][1]["content"][2]["text"] == "[MASKED]"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_masking_slice_stays_aligned_with_midturn_system(self):
|
||||
from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
|
||||
BedrockGuardrail,
|
||||
)
|
||||
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockMaskingGuardrail()
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"system": "trusted top-level system prompt",
|
||||
"messages": [
|
||||
{"role": "user", "content": "safe text"},
|
||||
{
|
||||
"role": "system",
|
||||
"content": [
|
||||
{"type": "text", "text": "prohibited correction"},
|
||||
{"type": "text", "text": "second correction"},
|
||||
],
|
||||
},
|
||||
{"role": "user", "content": "latest question"},
|
||||
],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert guardrail.inputs is not None
|
||||
texts = guardrail.inputs["texts"]
|
||||
structured = guardrail.inputs["structured_messages"]
|
||||
|
||||
bedrock = BedrockGuardrail(guardrailIdentifier="gi", guardrailVersion="1")
|
||||
total = sum(bedrock._count_message_texts(m) for m in structured)
|
||||
assert total == len(texts)
|
||||
|
||||
latest_user_index = bedrock._find_latest_message_index(structured, target_role="user")
|
||||
assert latest_user_index == 2
|
||||
scanned_slice = bedrock._locate_message_texts_slice(
|
||||
structured_messages=structured,
|
||||
target_index=latest_user_index,
|
||||
texts=texts,
|
||||
)
|
||||
assert scanned_slice == (3, 1)
|
||||
|
||||
merged = bedrock._merge_masked_texts(
|
||||
masked_texts=["{MASKED}"],
|
||||
texts=texts,
|
||||
scanned_slice=scanned_slice,
|
||||
scanned_role_subset=True,
|
||||
)
|
||||
assert merged == [
|
||||
"safe text",
|
||||
"prohibited correction",
|
||||
"second correction",
|
||||
"{MASKED}",
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_compaction_rewrite_keeps_midturn_system_messages(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockCompactingGuardrail(
|
||||
replacement_messages=[
|
||||
{"role": "user", "content": "compacted history"},
|
||||
{
|
||||
"role": "system",
|
||||
"content": [{"type": "text", "text": "use the corrected result"}],
|
||||
},
|
||||
{"role": "user", "content": "continue"},
|
||||
]
|
||||
)
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"system": "trusted top-level system prompt",
|
||||
"messages": [
|
||||
{"role": "user", "content": "original history"},
|
||||
{"role": "system", "content": "use the corrected result"},
|
||||
{"role": "user", "content": "continue"},
|
||||
],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert [m["role"] for m in data["messages"]] == ["user", "system", "user"]
|
||||
assert data["messages"][1]["content"] == [{"type": "text", "text": "use the corrected result"}]
|
||||
assert data["messages"][0]["content"] == [{"type": "text", "text": "compacted history"}]
|
||||
assert data["messages"][2]["content"] == [{"type": "text", "text": "continue"}]
|
||||
assert data["system"] == "trusted top-level system prompt"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_midturn_system_inside_tool_exchange_keeps_the_pair_intact(self):
|
||||
"""A system row between an assistant tool call and its result must not split the
|
||||
exchange into orphaned halves; it is emitted right after the exchange instead."""
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockCompactingGuardrail(
|
||||
replacement_messages=[
|
||||
{"role": "user", "content": "run the tool"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
"function": {"name": "get_weather", "arguments": "{}"},
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "system", "content": "use the corrected result"},
|
||||
{"role": "tool", "tool_call_id": "call_1", "content": "sunny"},
|
||||
]
|
||||
)
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"messages": [
|
||||
{"role": "user", "content": "run the tool"},
|
||||
{"role": "system", "content": "use the corrected result"},
|
||||
],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert [m["role"] for m in data["messages"]] == ["user", "assistant", "user", "system"]
|
||||
assistant_blocks = data["messages"][1]["content"]
|
||||
assert any(block.get("type") == "tool_use" and block.get("id") == "call_1" for block in assistant_blocks)
|
||||
result_blocks = data["messages"][2]["content"]
|
||||
assert [block["type"] for block in result_blocks] == ["tool_result"]
|
||||
assert result_blocks[0]["tool_use_id"] == "call_1"
|
||||
assert data["messages"][3]["content"] == "use the corrected result"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_compaction_rewrite_does_not_duplicate_hoisted_top_level_system(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockCompactingGuardrail(
|
||||
replacement_messages=[
|
||||
{"role": "system", "content": "trusted top-level system prompt"},
|
||||
{"role": "user", "content": "compacted history"},
|
||||
{"role": "system", "content": "use the corrected result"},
|
||||
]
|
||||
)
|
||||
guardrail.skip_system_message_in_guardrail = None
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"system": "trusted top-level system prompt",
|
||||
"messages": [
|
||||
{"role": "user", "content": "original history"},
|
||||
{"role": "system", "content": "use the corrected result"},
|
||||
],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert [m["role"] for m in data["messages"]] == ["user", "system"]
|
||||
assert data["messages"][1]["content"] == "use the corrected result"
|
||||
assert data["system"] == "trusted top-level system prompt"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_compaction_rewrite_keeps_leading_midturn_system_when_system_is_skipped(
|
||||
self,
|
||||
):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockCompactingGuardrail(
|
||||
replacement_messages=[
|
||||
{"role": "system", "content": "use the corrected result"},
|
||||
{"role": "user", "content": "compacted history"},
|
||||
]
|
||||
)
|
||||
guardrail.skip_system_message_in_guardrail = True
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"system": "trusted top-level system prompt",
|
||||
"messages": [
|
||||
{"role": "system", "content": "use the corrected result"},
|
||||
{"role": "user", "content": "original history"},
|
||||
],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert [m["role"] for m in data["messages"]] == ["system", "user"]
|
||||
assert data["messages"][0]["content"] == "use the corrected result"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_compaction_rewrite_keeps_leading_correction_when_top_level_system_hoists_nothing(
|
||||
self,
|
||||
):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockCompactingGuardrail(
|
||||
replacement_messages=[
|
||||
{"role": "system", "content": "use the corrected result"},
|
||||
{"role": "user", "content": "compacted history"},
|
||||
]
|
||||
)
|
||||
guardrail.skip_system_message_in_guardrail = None
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"system": [{"type": "image", "source": {"type": "url", "url": "https://example.com/a.png"}}],
|
||||
"messages": [
|
||||
{"role": "system", "content": "use the corrected result"},
|
||||
{"role": "user", "content": "original history"},
|
||||
],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert [m["role"] for m in data["messages"]] == ["system", "user"]
|
||||
assert data["messages"][0]["content"] == "use the corrected result"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_compaction_rewrite_keeps_leading_correction_when_hoisted_prompt_is_dropped(
|
||||
self,
|
||||
):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockCompactingGuardrail(
|
||||
replacement_messages=[
|
||||
{"role": "system", "content": "CLIENT CORRECTION"},
|
||||
{"role": "user", "content": "compacted history"},
|
||||
]
|
||||
)
|
||||
guardrail.skip_system_message_in_guardrail = None
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"system": "TRUSTED",
|
||||
"messages": [
|
||||
{"role": "system", "content": "CLIENT CORRECTION"},
|
||||
{"role": "user", "content": "original history"},
|
||||
],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert guardrail.inputs is not None
|
||||
assert guardrail.inputs["structured_messages"][0] == {
|
||||
"role": "system",
|
||||
"content": "TRUSTED",
|
||||
}
|
||||
assert [m["role"] for m in data["messages"]] == ["system", "user"]
|
||||
assert data["messages"][0]["content"] == "CLIENT CORRECTION"
|
||||
assert data["system"] == "TRUSTED"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_compaction_rewrite_drops_hoisted_prompt_matched_by_content_copy(self):
|
||||
import json
|
||||
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockCompactingGuardrail(
|
||||
replacement_messages=[
|
||||
json.loads(json.dumps({"role": "system", "content": "TRUSTED"})),
|
||||
{"role": "user", "content": "compacted history"},
|
||||
{"role": "system", "content": "CLIENT CORRECTION"},
|
||||
]
|
||||
)
|
||||
guardrail.skip_system_message_in_guardrail = None
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"system": "TRUSTED",
|
||||
"messages": [
|
||||
{"role": "user", "content": "original history"},
|
||||
{"role": "system", "content": "CLIENT CORRECTION"},
|
||||
],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert [m["role"] for m in data["messages"]] == ["user", "system"]
|
||||
assert data["messages"][1]["content"] == "CLIENT CORRECTION"
|
||||
assert data["system"] == "TRUSTED"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_compaction_rewrite_preserves_cache_control_on_system_blocks(self):
|
||||
"""
|
||||
`cache_control` on an in-sequence system text block survives the write-back, and is
|
||||
copied rather than aliased into the guardrail's own returned list.
|
||||
"""
|
||||
handler = AnthropicMessagesHandler()
|
||||
source_cache_control = {"type": "ephemeral"}
|
||||
guardrail = MockCompactingGuardrail(
|
||||
replacement_messages=[
|
||||
{"role": "user", "content": "compacted history"},
|
||||
{
|
||||
"role": "system",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "use the corrected result",
|
||||
"cache_control": source_cache_control,
|
||||
}
|
||||
],
|
||||
},
|
||||
]
|
||||
)
|
||||
guardrail.skip_system_message_in_guardrail = True
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"messages": [
|
||||
{"role": "user", "content": "original history"},
|
||||
{"role": "system", "content": "use the corrected result"},
|
||||
],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert data["messages"][1]["content"] == [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "use the corrected result",
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
}
|
||||
]
|
||||
assert data["messages"][1]["content"][0]["cache_control"] is not source_cache_control
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_compaction_rewrite_rstrips_trailing_assistant_in_each_run(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockCompactingGuardrail(
|
||||
replacement_messages=[
|
||||
{"role": "user", "content": "compacted history"},
|
||||
{"role": "assistant", "content": "earlier "},
|
||||
{"role": "system", "content": "use the corrected result"},
|
||||
{"role": "user", "content": "continue"},
|
||||
{"role": "assistant", "content": "prefill "},
|
||||
]
|
||||
)
|
||||
guardrail.skip_system_message_in_guardrail = True
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"messages": [
|
||||
{"role": "user", "content": "original history"},
|
||||
{"role": "system", "content": "use the corrected result"},
|
||||
],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert [m["role"] for m in data["messages"]] == [
|
||||
"user",
|
||||
"assistant",
|
||||
"system",
|
||||
"user",
|
||||
"assistant",
|
||||
]
|
||||
assert data["messages"][1]["content"] == [{"type": "text", "text": "earlier"}]
|
||||
assert data["messages"][-1]["content"] == [{"type": "text", "text": "prefill"}]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_compaction_rewrite_drops_text_free_system_message(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockCompactingGuardrail(
|
||||
replacement_messages=[
|
||||
{"role": "user", "content": "compacted history"},
|
||||
{"role": "system", "content": [{"type": "text", "text": ""}]},
|
||||
{"role": "system", "content": ""},
|
||||
{"role": "user", "content": "continue"},
|
||||
]
|
||||
)
|
||||
guardrail.skip_system_message_in_guardrail = True
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"messages": [
|
||||
{"role": "user", "content": "original history"},
|
||||
{"role": "system", "content": "use the corrected result"},
|
||||
{"role": "user", "content": "continue"},
|
||||
],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert [m["role"] for m in data["messages"]] == ["user", "user"]
|
||||
assert data["messages"][0]["content"] == [{"type": "text", "text": "compacted history"}]
|
||||
assert data["messages"][1]["content"] == [{"type": "text", "text": "continue"}]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_noncanonical_system_role_casing_is_still_scanned(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockMaskingGuardrail()
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"messages": [
|
||||
{"role": "user", "content": "safe text"},
|
||||
{"role": "System", "content": "prohibited correction"},
|
||||
],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert guardrail.inputs is not None
|
||||
assert "prohibited correction" in guardrail.inputs["texts"]
|
||||
assert data["messages"][1]["content"] == "[MASKED]"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_midturn_system_keeps_tool_result_turns_aligned_for_masking(self):
|
||||
"""Tool-result texts are scanned (LIT-5251), so counts align and the latest-user
|
||||
masking slice is locatable; a mid-turn system entry only shifts it by its own text."""
|
||||
from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
|
||||
BedrockGuardrail,
|
||||
)
|
||||
|
||||
handler = AnthropicMessagesHandler()
|
||||
bedrock = BedrockGuardrail(guardrailIdentifier="gi", guardrailVersion="1")
|
||||
tool_loop = [
|
||||
{"role": "user", "content": "call the tool"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [{"type": "tool_use", "id": "tu_1", "name": "get", "input": {"a": 1}}],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_result",
|
||||
"tool_use_id": "tu_1",
|
||||
"content": [{"type": "text", "text": "tool output"}],
|
||||
}
|
||||
],
|
||||
},
|
||||
]
|
||||
|
||||
async def _slice_for(messages: list):
|
||||
guardrail = MockMaskingGuardrail()
|
||||
data = {"model": "claude-3-5-sonnet-20241022", "messages": messages}
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
assert guardrail.inputs is not None
|
||||
texts = guardrail.inputs["texts"]
|
||||
structured = guardrail.inputs["structured_messages"]
|
||||
target_index = bedrock._find_latest_message_index(structured, target_role="user")
|
||||
return (
|
||||
sum(bedrock._count_message_texts(m) for m in structured) - len(texts),
|
||||
bedrock._locate_message_texts_slice(
|
||||
structured_messages=structured,
|
||||
target_index=target_index,
|
||||
texts=texts,
|
||||
),
|
||||
)
|
||||
|
||||
with_system = await _slice_for(
|
||||
tool_loop
|
||||
+ [
|
||||
{"role": "system", "content": "use the corrected result"},
|
||||
{"role": "user", "content": "latest question"},
|
||||
]
|
||||
)
|
||||
without_system = await _slice_for(tool_loop + [{"role": "user", "content": "latest question"}])
|
||||
|
||||
assert with_system == (0, (3, 1))
|
||||
assert without_system == (0, (2, 1))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_compaction_rewrite_to_only_system_messages_is_rejected(self):
|
||||
import litellm
|
||||
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockCompactingGuardrail(
|
||||
replacement_messages=[{"role": "system", "content": "use the corrected result"}]
|
||||
)
|
||||
guardrail.skip_system_message_in_guardrail = True
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"messages": [
|
||||
{"role": "user", "content": "original history"},
|
||||
{"role": "system", "content": "use the corrected result"},
|
||||
],
|
||||
}
|
||||
|
||||
with patch.object(litellm, "modify_params", False):
|
||||
with pytest.raises(litellm.BadRequestError, match="at least one non-system message"):
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_compaction_rewrite_to_only_system_messages_repaired_with_modify_params(
|
||||
self,
|
||||
):
|
||||
import litellm
|
||||
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockCompactingGuardrail(
|
||||
replacement_messages=[{"role": "system", "content": "use the corrected result"}]
|
||||
)
|
||||
guardrail.skip_system_message_in_guardrail = True
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"messages": [
|
||||
{"role": "user", "content": "original history"},
|
||||
{"role": "system", "content": "use the corrected result"},
|
||||
],
|
||||
}
|
||||
|
||||
with patch.object(litellm, "modify_params", True):
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert [m["role"] for m in data["messages"]] == ["system", "user"]
|
||||
assert data["messages"][0]["content"] == "use the corrected result"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_compaction_rewrite_without_system_messages_is_unchanged(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockCompactingGuardrail(replacement_messages=[{"role": "user", "content": "compacted history"}])
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"messages": [
|
||||
{"role": "user", "content": "a"},
|
||||
{"role": "assistant", "content": "b"},
|
||||
{"role": "user", "content": "c"},
|
||||
],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert data["messages"] == [{"role": "user", "content": [{"type": "text", "text": "compacted history"}]}]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_output_streaming_response_empty_choices(self):
|
||||
"""Test that streaming response with empty choices doesn't raise IndexError
|
||||
|
|
@ -597,7 +1340,7 @@ class TestAnthropicMessagesIncrementalScan:
|
|||
assert "Thanks, summarize the result." in scanned
|
||||
|
||||
|
||||
class MockMaskingGuardrail(CustomGuardrail):
|
||||
class MockCanaryMaskingGuardrail(CustomGuardrail):
|
||||
"""Records every text handed to it and masks a canary token in place."""
|
||||
|
||||
def __init__(self, guardrail_name: str = "mask-canary"):
|
||||
|
|
@ -629,7 +1372,7 @@ class TestAnthropicMessagesToolResultScanning:
|
|||
@pytest.mark.asyncio
|
||||
async def test_string_form_tool_result_is_scanned_and_written_back(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockMaskingGuardrail()
|
||||
guardrail = MockCanaryMaskingGuardrail()
|
||||
messages = [
|
||||
{"role": "user", "content": "fetch the page"},
|
||||
{
|
||||
|
|
@ -652,7 +1395,7 @@ class TestAnthropicMessagesToolResultScanning:
|
|||
@pytest.mark.asyncio
|
||||
async def test_list_form_tool_result_is_scanned_and_written_back(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockMaskingGuardrail()
|
||||
guardrail = MockCanaryMaskingGuardrail()
|
||||
messages = [
|
||||
{"role": "user", "content": "fetch the page"},
|
||||
{
|
||||
|
|
@ -683,7 +1426,7 @@ class TestAnthropicMessagesToolResultScanning:
|
|||
"""The write-back is positional, so a single mis-indexed target silently
|
||||
writes one message's masked text over another's."""
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockMaskingGuardrail()
|
||||
guardrail = MockCanaryMaskingGuardrail()
|
||||
messages = [
|
||||
{"role": "user", "content": "plain POISON string"},
|
||||
{
|
||||
|
|
@ -713,7 +1456,7 @@ class TestAnthropicMessagesToolResultScanning:
|
|||
async def test_image_inside_tool_result_is_collected(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
|
||||
class ImageRecordingGuardrail(MockMaskingGuardrail):
|
||||
class ImageRecordingGuardrail(MockCanaryMaskingGuardrail):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.seen_images: list[str] = []
|
||||
|
|
@ -746,7 +1489,7 @@ class TestAnthropicMessagesToolResultScanning:
|
|||
@pytest.mark.asyncio
|
||||
async def test_tool_result_is_skipped_when_guardrail_skips_tool_messages(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockMaskingGuardrail()
|
||||
guardrail = MockCanaryMaskingGuardrail()
|
||||
guardrail.skip_tool_message_in_guardrail = True
|
||||
messages = [
|
||||
{"role": "user", "content": "keep me POISON"},
|
||||
|
|
@ -763,7 +1506,7 @@ class TestAnthropicMessagesToolResultScanning:
|
|||
assert messages[0]["content"] == "keep me [BLOCKED]"
|
||||
|
||||
|
||||
class InputsRecordingGuardrail(MockMaskingGuardrail):
|
||||
class InputsRecordingGuardrail(MockCanaryMaskingGuardrail):
|
||||
def __init__(self):
|
||||
super().__init__(guardrail_name="scan-only-capture")
|
||||
self.captured_inputs: Optional[GenericGuardrailAPIInputs] = None
|
||||
|
|
|
|||
|
|
@ -413,6 +413,224 @@ def test_translate_anthropic_messages_to_openai_tool_message_placement():
|
|||
), "Tool message should be placed before user message"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("system_content", "expected_content"),
|
||||
[
|
||||
("Use the corrected result.", "Use the corrected result."),
|
||||
(
|
||||
[{"type": "text", "text": "Use the corrected result."}],
|
||||
[{"type": "text", "text": "Use the corrected result."}],
|
||||
),
|
||||
(
|
||||
[
|
||||
{
|
||||
"type": "image",
|
||||
"source": {"type": "url", "url": "https://example.com/a.png"},
|
||||
},
|
||||
{"type": "text", "text": "Use the corrected result."},
|
||||
],
|
||||
[{"type": "text", "text": "Use the corrected result."}],
|
||||
),
|
||||
(
|
||||
[
|
||||
{"type": "text", "text": "First correction."},
|
||||
{"type": "text", "text": "Second correction."},
|
||||
],
|
||||
[
|
||||
{"type": "text", "text": "First correction."},
|
||||
{"type": "text", "text": "Second correction."},
|
||||
],
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_translate_anthropic_messages_to_openai_preserves_midturn_system_correction(
|
||||
system_content: object,
|
||||
expected_content: object,
|
||||
):
|
||||
messages = [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_use",
|
||||
"id": "toolu_01234",
|
||||
"name": "get_weather",
|
||||
"input": {"location": "Boston"},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_result",
|
||||
"tool_use_id": "toolu_01234",
|
||||
"content": "Rainy, 55°F",
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "system", "content": system_content},
|
||||
{"role": "user", "content": "Continue."},
|
||||
]
|
||||
|
||||
result = LiteLLMAnthropicMessagesAdapter().translate_anthropic_messages_to_openai(
|
||||
messages=messages,
|
||||
model="claude-3-5-sonnet-20240620",
|
||||
)
|
||||
|
||||
assert result == [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"thinking_blocks": None,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "toolu_01234",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"arguments": '{"location": "Boston"}',
|
||||
},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "toolu_01234",
|
||||
"content": "Rainy, 55°F",
|
||||
},
|
||||
{"role": "system", "content": expected_content},
|
||||
{"role": "user", "content": "Continue."},
|
||||
]
|
||||
|
||||
|
||||
def test_translate_anthropic_messages_to_openai_preserves_midturn_system_cache_control():
|
||||
"""
|
||||
`cache_control` on an in-sequence system text block survives, matching how the
|
||||
hoisted top-level `system` prompt and user text blocks are already handled.
|
||||
"""
|
||||
messages = [
|
||||
{
|
||||
"role": "system",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Use the corrected result.",
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
}
|
||||
],
|
||||
}
|
||||
]
|
||||
|
||||
result = LiteLLMAnthropicMessagesAdapter().translate_anthropic_messages_to_openai(
|
||||
messages=messages,
|
||||
model="claude-3-5-sonnet-20240620",
|
||||
)
|
||||
|
||||
assert result == [
|
||||
{
|
||||
"role": "system",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Use the corrected result.",
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
}
|
||||
],
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
def test_translate_anthropic_messages_to_openai_drops_midturn_system_cache_control_for_non_claude():
|
||||
"""
|
||||
`cache_control` goes through the same `_add_cache_control_if_applicable` gate as the
|
||||
hoisted top-level prompt and user text blocks, so a non-Claude *requested model name*
|
||||
does not get it. That gate is a best-effort check of the requested name before routing
|
||||
(behind the proxy it is often a public alias), not a guarantee about the backend that
|
||||
ultimately serves the request.
|
||||
"""
|
||||
messages = [
|
||||
{
|
||||
"role": "system",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Use the corrected result.",
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
}
|
||||
],
|
||||
}
|
||||
]
|
||||
|
||||
result = LiteLLMAnthropicMessagesAdapter().translate_anthropic_messages_to_openai(
|
||||
messages=messages,
|
||||
model="gpt-4o",
|
||||
)
|
||||
|
||||
assert result == [
|
||||
{
|
||||
"role": "system",
|
||||
"content": [{"type": "text", "text": "Use the corrected result."}],
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"system_content",
|
||||
[
|
||||
"",
|
||||
[{"type": "text", "text": ""}],
|
||||
[
|
||||
{
|
||||
"type": "image",
|
||||
"source": {"type": "url", "url": "https://example.com/a.png"},
|
||||
}
|
||||
],
|
||||
None,
|
||||
],
|
||||
)
|
||||
def test_translate_anthropic_messages_to_openai_drops_empty_midturn_system(
|
||||
system_content: object,
|
||||
):
|
||||
messages = [{"role": "system", "content": system_content}]
|
||||
|
||||
result = LiteLLMAnthropicMessagesAdapter().translate_anthropic_messages_to_openai(
|
||||
messages=messages,
|
||||
model="claude-3-5-sonnet-20240620",
|
||||
)
|
||||
|
||||
assert result == []
|
||||
|
||||
|
||||
def test_translate_anthropic_to_openai_orders_top_level_and_midturn_system():
|
||||
"""
|
||||
Request level: the trusted top-level prompt is hoisted to index 0 exactly once and the
|
||||
in-sequence correction keeps its own position and `role: "system"` -- no duplication of
|
||||
either, and no reordering of the surrounding turns.
|
||||
"""
|
||||
openai_request, _ = LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai(
|
||||
anthropic_message_request={
|
||||
"model": "claude-3-5-sonnet-20240620",
|
||||
"max_tokens": 100,
|
||||
"system": "Trusted top-level prompt.",
|
||||
"messages": [
|
||||
{"role": "user", "content": "First question."},
|
||||
{"role": "assistant", "content": "First answer."},
|
||||
{"role": "system", "content": "Use the corrected result."},
|
||||
{"role": "user", "content": "Continue."},
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
assert openai_request["messages"] == [
|
||||
{"role": "system", "content": "Trusted top-level prompt."},
|
||||
{"role": "user", "content": "First question."},
|
||||
{"role": "assistant", "content": "First answer.", "thinking_blocks": None},
|
||||
{"role": "system", "content": "Use the corrected result."},
|
||||
{"role": "user", "content": "Continue."},
|
||||
]
|
||||
|
||||
|
||||
def test_translate_openai_content_to_anthropic_empty_function_arguments():
|
||||
"""Test that empty function arguments are handled safely and don't cause JSON parsing errors."""
|
||||
|
||||
|
|
|
|||
|
|
@ -2475,3 +2475,57 @@ def test_endpoint_runs_failure_hook_on_500_context_management_error():
|
|||
body = response.json()
|
||||
assert body["type"] == "error"
|
||||
failure_hook.assert_awaited_once()
|
||||
|
||||
|
||||
def test_count_effective_tokens_counts_midturn_system_correction():
|
||||
from litellm.llms.anthropic.experimental_pass_through.context_management.editors.compact import (
|
||||
_count_effective_tokens,
|
||||
)
|
||||
|
||||
base: List[Dict[str, Any]] = [
|
||||
{"role": "user", "content": "hello"},
|
||||
{"role": "assistant", "content": "hi"},
|
||||
]
|
||||
correction = {
|
||||
"role": "system",
|
||||
"content": [{"type": "text", "text": "use the corrected result " * 20}],
|
||||
}
|
||||
|
||||
without_correction = _count_effective_tokens(
|
||||
model=MODEL, effective_messages=base, compaction_block=None, tools=None
|
||||
)
|
||||
with_correction = _count_effective_tokens(
|
||||
model=MODEL,
|
||||
effective_messages=base + [correction],
|
||||
compaction_block=None,
|
||||
tools=None,
|
||||
)
|
||||
|
||||
assert with_correction > without_correction
|
||||
|
||||
|
||||
def test_build_summary_messages_keeps_midturn_system_correction_in_place():
|
||||
from litellm.llms.anthropic.experimental_pass_through.context_management.editors.compact import (
|
||||
_build_summary_messages,
|
||||
)
|
||||
|
||||
summary_messages = _build_summary_messages(
|
||||
effective_messages=[
|
||||
{"role": "user", "content": "original question"},
|
||||
{"role": "system", "content": "use the corrected result"},
|
||||
{"role": "assistant", "content": "acknowledged"},
|
||||
],
|
||||
prompt="summarize the conversation",
|
||||
system="caller system prompt",
|
||||
)
|
||||
|
||||
assert [m["role"] for m in summary_messages] == [
|
||||
"system",
|
||||
"user",
|
||||
"system",
|
||||
"assistant",
|
||||
"user",
|
||||
]
|
||||
assert summary_messages[0]["content"] == "caller system prompt"
|
||||
assert summary_messages[2]["content"] == "use the corrected result"
|
||||
assert summary_messages[-1]["content"] == "summarize the conversation"
|
||||
|
|
|
|||
|
|
@ -9,6 +9,8 @@ import sys
|
|||
from typing import Any, Dict, List
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../../../../../.."))
|
||||
|
||||
from litellm.constants import (
|
||||
|
|
@ -222,6 +224,106 @@ class TestTranslateMessagesToResponsesInput:
|
|||
{"type": "input_text", "text": "Second part."},
|
||||
]
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"system_content",
|
||||
[
|
||||
"Use the corrected result.",
|
||||
[{"type": "text", "text": "Use the corrected result."}],
|
||||
[
|
||||
{"type": "image", "source": {"type": "url", "url": "https://example.com/a.png"}},
|
||||
{"type": "text", "text": "Use the corrected result."},
|
||||
],
|
||||
],
|
||||
)
|
||||
def test_midturn_system_correction_stays_system_in_sequence(self, system_content: object):
|
||||
messages = [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_use",
|
||||
"id": "toolu_01234",
|
||||
"name": "get_weather",
|
||||
"input": {"location": "Boston"},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_result",
|
||||
"tool_use_id": "toolu_01234",
|
||||
"content": "Rainy, 55°F",
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "system", "content": system_content},
|
||||
{"role": "user", "content": "Continue."},
|
||||
]
|
||||
|
||||
result = _translate_messages(messages)
|
||||
|
||||
assert result == [
|
||||
{
|
||||
"type": "function_call",
|
||||
"call_id": "toolu_01234",
|
||||
"name": "get_weather",
|
||||
"arguments": '{"location": "Boston"}',
|
||||
},
|
||||
{
|
||||
"type": "function_call_output",
|
||||
"call_id": "toolu_01234",
|
||||
"output": "Rainy, 55°F",
|
||||
},
|
||||
{
|
||||
"type": "message",
|
||||
"role": "system",
|
||||
"content": [{"type": "input_text", "text": "Use the corrected result."}],
|
||||
},
|
||||
{
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": [{"type": "input_text", "text": "Continue."}],
|
||||
},
|
||||
]
|
||||
|
||||
def test_midturn_system_correction_keeps_multiple_text_blocks(self):
|
||||
messages = [
|
||||
{
|
||||
"role": "system",
|
||||
"content": [
|
||||
{"type": "text", "text": "First correction."},
|
||||
{"type": "text", "text": "Second correction."},
|
||||
],
|
||||
}
|
||||
]
|
||||
|
||||
assert _translate_messages(messages) == [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "system",
|
||||
"content": [
|
||||
{"type": "input_text", "text": "First correction."},
|
||||
{"type": "input_text", "text": "Second correction."},
|
||||
],
|
||||
}
|
||||
]
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"system_content",
|
||||
[
|
||||
"",
|
||||
[{"type": "text", "text": ""}],
|
||||
[{"type": "image", "source": {"type": "url", "url": "https://example.com/a.png"}}],
|
||||
None,
|
||||
],
|
||||
)
|
||||
def test_empty_or_unsupported_midturn_system_correction_is_dropped(self, system_content: object):
|
||||
messages = [{"role": "system", "content": system_content}]
|
||||
|
||||
assert _translate_messages(messages) == []
|
||||
|
||||
def test_user_base64_image(self):
|
||||
"""User message with base64 image source becomes input_image with data URL."""
|
||||
messages = [
|
||||
|
|
@ -723,6 +825,42 @@ class TestTranslateRequestBroaderCoverage:
|
|||
kwargs = _ADAPTER.translate_request(req)
|
||||
assert kwargs["instructions"] == "You are a helpful assistant."
|
||||
|
||||
def test_top_level_system_and_midturn_correction_are_not_duplicated(self):
|
||||
"""
|
||||
Request level: the trusted top-level prompt goes to `instructions` only, and the
|
||||
in-sequence correction stays a `role: "system"` input item in its original position.
|
||||
Neither appears twice, and the surrounding turns keep their order.
|
||||
"""
|
||||
req = _make_request(
|
||||
system="Trusted top-level prompt.",
|
||||
messages=[
|
||||
{"role": "user", "content": "First question."},
|
||||
{"role": "system", "content": "Use the corrected result."},
|
||||
{"role": "user", "content": "Continue."},
|
||||
],
|
||||
)
|
||||
|
||||
kwargs = _ADAPTER.translate_request(req)
|
||||
|
||||
assert kwargs["instructions"] == "Trusted top-level prompt."
|
||||
assert kwargs["input"] == [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": [{"type": "input_text", "text": "First question."}],
|
||||
},
|
||||
{
|
||||
"type": "message",
|
||||
"role": "system",
|
||||
"content": [{"type": "input_text", "text": "Use the corrected result."}],
|
||||
},
|
||||
{
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": [{"type": "input_text", "text": "Continue."}],
|
||||
},
|
||||
]
|
||||
|
||||
def test_system_list_of_text_blocks_joined(self):
|
||||
req = _make_request(
|
||||
system=[
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue