mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
fix(policy_engine): snapshot guardrail inputs before apply_guardrail so in-place stream rewrites are withheld
This commit is contained in:
parent
45a6b1de23
commit
badefa395c
2 changed files with 34 additions and 10 deletions
|
|
@ -57,16 +57,16 @@ def _tool_call_shape(tool_call: object) -> tuple[object, object]:
|
|||
return (function.get("name"), function.get("arguments"))
|
||||
|
||||
|
||||
def _rewrote_texts(sent: Sequence[str] | None, returned: Sequence[str] | None) -> bool:
|
||||
return sent is not None and returned is not None and tuple(returned) != tuple(sent)
|
||||
def _text_snapshot(texts: Sequence[str] | None) -> tuple[str, ...] | None:
|
||||
return None if texts is None else tuple(texts)
|
||||
|
||||
|
||||
def _rewrote_tool_calls(sent: Sequence[object] | None, returned: Sequence[object] | None) -> bool:
|
||||
if sent is None or returned is None:
|
||||
return False
|
||||
return tuple(_tool_call_shape(tool_call) for tool_call in returned) != tuple(
|
||||
_tool_call_shape(tool_call) for tool_call in sent
|
||||
)
|
||||
def _tool_call_shapes(tool_calls: Sequence[object] | None) -> tuple[tuple[object, object], ...] | None:
|
||||
return None if tool_calls is None else tuple(_tool_call_shape(tool_call) for tool_call in tool_calls)
|
||||
|
||||
|
||||
def _rewrote(sent: tuple[object, ...] | None, returned: tuple[object, ...] | None) -> bool:
|
||||
return sent is not None and returned is not None and returned != sent
|
||||
|
||||
|
||||
class _StreamRewriteObserver(CustomGuardrail):
|
||||
|
|
@ -90,13 +90,15 @@ class _StreamRewriteObserver(CustomGuardrail):
|
|||
input_type: Literal["request", "response"],
|
||||
logging_obj: "LiteLLMLoggingObj | None" = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
sent_texts: Final = _text_snapshot(inputs.get("texts"))
|
||||
sent_tool_shapes: Final = _tool_call_shapes(inputs.get("tool_calls"))
|
||||
outputs: Final = await self.inner.apply_guardrail(
|
||||
inputs=inputs, request_data=request_data, input_type=input_type, logging_obj=logging_obj
|
||||
)
|
||||
self.rewrote = (
|
||||
self.rewrote
|
||||
or _rewrote_texts(inputs.get("texts"), outputs.get("texts"))
|
||||
or _rewrote_tool_calls(inputs.get("tool_calls"), outputs.get("tool_calls"))
|
||||
or _rewrote(sent_texts, _text_snapshot(outputs.get("texts")))
|
||||
or _rewrote(sent_tool_shapes, _tool_call_shapes(outputs.get("tool_calls")))
|
||||
)
|
||||
return outputs
|
||||
|
||||
|
|
|
|||
|
|
@ -872,3 +872,25 @@ async def test_streaming_step_unchanged_texts_in_another_container_allow(monkeyp
|
|||
|
||||
assert result.terminal_action == "allow"
|
||||
assert [step.outcome for step in result.step_results] == ["pass"]
|
||||
|
||||
|
||||
class _InPlaceMutatingGuardrail(CustomGuardrail):
|
||||
"""Rewrites like bedrock/presidio do: rebinds inputs["texts"] on the dict it was handed
|
||||
and returns that same dict, so a post-call comparison against inputs sees no change."""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(guardrail_name="masker", event_hook="post_call", default_on=True)
|
||||
|
||||
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
|
||||
inputs["texts"] = ["hello [MASKED]"]
|
||||
return inputs
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_step_in_place_rewrite_still_withholds_stream(monkeypatch):
|
||||
monkeypatch.setattr(litellm, "callbacks", [_InPlaceMutatingGuardrail()])
|
||||
|
||||
with pytest.raises(UndeliverableStreamRewrite) as info:
|
||||
await _run_streaming_step(["hello [MASKED]"], _TextTranslation())
|
||||
|
||||
assert info.value.guardrail_name == "masker"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue