mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge pull request #37231 from BerriAI/litellm_lit5696_system_hoist_writeback
fix(anthropic): fold guardrail-modified leading system rows into top-level system param
This commit is contained in:
commit
e81cedb13a
2 changed files with 179 additions and 14 deletions
|
|
@ -499,6 +499,39 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
{"role": "system", "content": blocks} if blocks else None # mutable-ok: API message payload
|
||||
) # mutable-ok: API message payload
|
||||
|
||||
@staticmethod
|
||||
def _fold_leading_systems_into_top_level(
|
||||
data: dict[str, object], # mutable-ok: API message payload
|
||||
leading_systems: Sequence[object],
|
||||
include_existing_system: bool,
|
||||
) -> None:
|
||||
"""Deliver leading system rows through Anthropic's top-level system param, which rejects them in messages."""
|
||||
existing: Final = data.get("system") if include_existing_system else None
|
||||
existing_blocks: Final[list[object]] = ( # mutable-ok: API message payload
|
||||
[{"type": "text", "text": existing}]
|
||||
if isinstance(existing, str) and existing
|
||||
else list(existing)
|
||||
if isinstance(existing, list)
|
||||
else []
|
||||
)
|
||||
converted_rows: Final = tuple(
|
||||
AnthropicMessagesHandler._openai_system_message_to_anthropic(message)
|
||||
for message in leading_systems
|
||||
if isinstance(message, dict)
|
||||
)
|
||||
folded: Final[list[object]] = existing_blocks + [ # mutable-ok: API message payload
|
||||
block
|
||||
for row in converted_rows
|
||||
if row is not None
|
||||
for block in (
|
||||
[{"type": "text", "text": row["content"]}] if isinstance(row["content"], str) else row["content"]
|
||||
)
|
||||
]
|
||||
if folded:
|
||||
data["system"] = folded # rebind-ok: write-back mutates the request payload in place
|
||||
else:
|
||||
data.pop("system", None)
|
||||
|
||||
@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."""
|
||||
|
|
@ -575,9 +608,24 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
)
|
||||
|
||||
ordered: Final = AnthropicMessagesHandler._defer_systems_inside_tool_exchanges(structured_messages)
|
||||
leading_count: Final = next(
|
||||
(index for index, message in enumerate(ordered) if not _is_system(message)),
|
||||
len(ordered),
|
||||
)
|
||||
leading_systems: Final = ordered[:leading_count]
|
||||
hoisted_in_leading: Final = any(
|
||||
AnthropicMessagesHandler._is_hoisted_top_level_system(message, hoisted_system_message)
|
||||
for message in leading_systems
|
||||
)
|
||||
if leading_systems and not (leading_count == 1 and hoisted_in_leading):
|
||||
AnthropicMessagesHandler._fold_leading_systems_into_top_level(
|
||||
data,
|
||||
leading_systems,
|
||||
include_existing_system=hoisted_system_message is None,
|
||||
)
|
||||
run: Final[list] = [] # mutable-ok: API message payload
|
||||
hoisted_dropped = False # rebind-ok: flips once the hoisted prompt is dropped
|
||||
for message in ordered:
|
||||
hoisted_dropped = hoisted_in_leading # rebind-ok: flips once the hoisted prompt is dropped
|
||||
for message in ordered[leading_count:]:
|
||||
if not _is_system(message):
|
||||
run.append(message)
|
||||
continue
|
||||
|
|
|
|||
|
|
@ -121,6 +121,45 @@ class MockCompactingGuardrail(CustomGuardrail):
|
|||
return rewritten
|
||||
|
||||
|
||||
class MockStructuredMaskingGuardrail(CustomGuardrail):
|
||||
"""Mask an email in texts and in a rebuilt structured view, like a PII-masking guardrail (LIT-5696)."""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(guardrail_name="structured-masking-test")
|
||||
|
||||
@staticmethod
|
||||
def _mask(text: str) -> str:
|
||||
return text.replace("bob@example.com", "<EMAIL>")
|
||||
|
||||
def _mask_content(self, content: object) -> object:
|
||||
if isinstance(content, str):
|
||||
return self._mask(content)
|
||||
if not isinstance(content, list):
|
||||
return content
|
||||
return [
|
||||
{**block, "text": self._mask(block["text"])}
|
||||
if isinstance(block, dict) and isinstance(block.get("text"), str)
|
||||
else block
|
||||
for block in content
|
||||
]
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional[Any] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
masked = inputs.copy()
|
||||
masked["texts"] = [self._mask(text) for text in inputs.get("texts", [])]
|
||||
structured = inputs.get("structured_messages")
|
||||
if structured is not None:
|
||||
masked["structured_messages"] = [
|
||||
{**message, "content": self._mask_content(message.get("content"))} for message in structured
|
||||
]
|
||||
return masked
|
||||
|
||||
|
||||
class TestAnthropicMessagesHandlerStreamingRequestData:
|
||||
"""Post-call guardrails on streaming /v1/messages receive the response and identity metadata"""
|
||||
|
||||
|
|
@ -602,7 +641,7 @@ class TestAnthropicMessagesHandlerInputProcessing:
|
|||
assert data["system"] == "trusted top-level system prompt"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_compaction_rewrite_keeps_leading_midturn_system_when_system_is_skipped(
|
||||
async def test_leading_system_row_appends_to_skipped_top_level_system(
|
||||
self,
|
||||
):
|
||||
handler = AnthropicMessagesHandler()
|
||||
|
|
@ -624,11 +663,14 @@ class TestAnthropicMessagesHandlerInputProcessing:
|
|||
|
||||
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"
|
||||
assert [m["role"] for m in data["messages"]] == ["user"]
|
||||
assert data["system"] == [
|
||||
{"type": "text", "text": "trusted top-level system prompt"},
|
||||
{"type": "text", "text": "use the corrected result"},
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_compaction_rewrite_keeps_leading_correction_when_top_level_system_hoists_nothing(
|
||||
async def test_leading_correction_appends_when_top_level_system_hoists_nothing(
|
||||
self,
|
||||
):
|
||||
handler = AnthropicMessagesHandler()
|
||||
|
|
@ -650,11 +692,14 @@ class TestAnthropicMessagesHandlerInputProcessing:
|
|||
|
||||
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"
|
||||
assert [m["role"] for m in data["messages"]] == ["user"]
|
||||
assert data["system"] == [
|
||||
{"type": "image", "source": {"type": "url", "url": "https://example.com/a.png"}},
|
||||
{"type": "text", "text": "use the corrected result"},
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_compaction_rewrite_keeps_leading_correction_when_hoisted_prompt_is_dropped(
|
||||
async def test_leading_correction_replaces_top_level_system_when_hoisted_prompt_is_dropped(
|
||||
self,
|
||||
):
|
||||
handler = AnthropicMessagesHandler()
|
||||
|
|
@ -681,9 +726,81 @@ class TestAnthropicMessagesHandlerInputProcessing:
|
|||
"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"
|
||||
assert [m["role"] for m in data["messages"]] == ["user"]
|
||||
assert data["system"] == [{"type": "text", "text": "CLIENT CORRECTION"}]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_masked_hoisted_system_folds_into_top_level_system(self):
|
||||
"""LIT-5696: a guardrail-modified top-level prompt must go back through the system
|
||||
param; emitting it as messages[0] is rejected by Anthropic, dropping it leaks the
|
||||
unmasked original."""
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockStructuredMaskingGuardrail()
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"system": [{"type": "text", "text": "You are helpful. The admin is bob@example.com."}],
|
||||
"messages": [{"role": "user", "content": [{"type": "text", "text": "hi"}]}],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert data["system"] == [{"type": "text", "text": "You are helpful. The admin is <EMAIL>."}]
|
||||
assert [m["role"] for m in data["messages"]] == ["user"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_client_leading_system_row_folds_into_top_level_system(self):
|
||||
"""LIT-5696: a client-sent leading system row folds into the system param instead of
|
||||
being sent back as messages[0], which Anthropic rejects."""
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockStructuredMaskingGuardrail()
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"messages": [
|
||||
{"role": "system", "content": [{"type": "text", "text": "You are helpful."}]},
|
||||
{"role": "user", "content": [{"type": "text", "text": "hi"}]},
|
||||
],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert data["system"] == [{"type": "text", "text": "You are helpful."}]
|
||||
assert [m["role"] for m in data["messages"]] == ["user"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_masked_midturn_system_after_user_stays_in_messages(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockStructuredMaskingGuardrail()
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"system": [{"type": "text", "text": "You are helpful. The admin is bob@example.com."}],
|
||||
"messages": [
|
||||
{"role": "user", "content": [{"type": "text", "text": "hi"}]},
|
||||
{"role": "assistant", "content": [{"type": "text", "text": "hello"}]},
|
||||
{"role": "system", "content": [{"type": "text", "text": "Mid-turn: admin bob@example.com"}]},
|
||||
{"role": "user", "content": [{"type": "text", "text": "next"}]},
|
||||
],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert data["system"] == [{"type": "text", "text": "You are helpful. The admin is <EMAIL>."}]
|
||||
assert [m["role"] for m in data["messages"]] == ["user", "assistant", "system", "user"]
|
||||
assert data["messages"][2]["content"] == [{"type": "text", "text": "Mid-turn: admin <EMAIL>"}]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unmodified_structured_copy_leaves_top_level_system_untouched(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockStructuredMaskingGuardrail()
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"system": "You are helpful.",
|
||||
"messages": [{"role": "user", "content": [{"type": "text", "text": "hi"}]}],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert data["system"] == "You are helpful."
|
||||
assert [m["role"] for m in data["messages"]] == ["user"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_compaction_rewrite_drops_hoisted_prompt_matched_by_content_copy(self):
|
||||
|
|
@ -934,8 +1051,8 @@ class TestAnthropicMessagesHandlerInputProcessing:
|
|||
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"
|
||||
assert data["messages"] == [{"role": "user", "content": [{"type": "text", "text": "Please continue."}]}]
|
||||
assert data["system"] == [{"type": "text", "text": "use the corrected result"}]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_compaction_rewrite_without_system_messages_is_unchanged(self):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue