mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(guardrails): preserve tool context during streamed text scans
This commit is contained in:
parent
cb01259970
commit
400f845b5f
2 changed files with 19 additions and 4 deletions
|
|
@ -794,6 +794,14 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
self.merge_user_api_key_metadata_into_request(request_data, user_api_key_dict)
|
||||
|
||||
inputs: Final = GenericGuardrailAPIInputs(texts=texts_to_check)
|
||||
if self._streamed_tool_call_fingerprints(responses_so_far):
|
||||
assembled: Final = self._rebuild_ended_stream_per_choice(responses_so_far, litellm_logging_obj)
|
||||
inputs["tool_calls"] = [
|
||||
converted
|
||||
for choice in assembled.choices
|
||||
for tool_call in choice.message.tool_calls or ()
|
||||
if (converted := self._convert_tool_call_to_dict(tool_call)) is not None
|
||||
]
|
||||
if responses_so_far and getattr(responses_so_far[0], "model", None):
|
||||
inputs["model"] = responses_so_far[0].model
|
||||
guardrailed_inputs: Final = await guardrail_to_apply.apply_guardrail(
|
||||
|
|
|
|||
|
|
@ -946,10 +946,11 @@ def _delta_text(item):
|
|||
|
||||
|
||||
class _ToolRedactingGuardrail(CustomGuardrail):
|
||||
def __init__(self) -> None:
|
||||
def __init__(self, require_tool_context: bool = False) -> None:
|
||||
super().__init__(guardrail_name="tool-redactor", event_hook=GuardrailEventHooks.post_call, default_on=True)
|
||||
self.streaming_transform_mode = "incremental_diff"
|
||||
self.streaming_sampling_rate = 1
|
||||
self.require_tool_context = require_tool_context
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
|
|
@ -961,7 +962,10 @@ class _ToolRedactingGuardrail(CustomGuardrail):
|
|||
calls: Final = TypeAdapter(tuple[ChatCompletionMessageToolCall, ...]).validate_python(
|
||||
inputs.get("tool_calls", ())
|
||||
)
|
||||
texts: Final = tuple("checked:" + text.replace("SECRET", "MASKED") for text in inputs.get("texts", ()))
|
||||
texts: Final = tuple(
|
||||
"checked:" + text.replace("SECRET", "MASKED") if calls or not self.require_tool_context else text
|
||||
for text in inputs.get("texts", ())
|
||||
)
|
||||
return {
|
||||
**inputs,
|
||||
"texts": list(texts),
|
||||
|
|
@ -984,10 +988,11 @@ class TestStreamingTransform:
|
|||
_patch_translation_mappings(monkeypatch, {CallTypes.acompletion: OpenAIChatCompletionsHandler})
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("require_tool_context", [False, True])
|
||||
@pytest.mark.parametrize("include_text", [False, True])
|
||||
@pytest.mark.parametrize("tool_count", [1, 2])
|
||||
async def test_buffered_tool_arguments_are_rewritten_before_delivery(
|
||||
self, include_text: bool, tool_count: int
|
||||
self, include_text: bool, tool_count: int, require_tool_context: bool
|
||||
) -> None:
|
||||
chunks: Final = (
|
||||
*([_stream_chunk("hello SECRET")] if include_text else []),
|
||||
|
|
@ -1001,7 +1006,9 @@ class TestStreamingTransform:
|
|||
]), finish_reason="tool_calls")]),
|
||||
ModelResponseStream(choices=[], usage={"prompt_tokens": 7, "completion_tokens": 11, "total_tokens": 18}),
|
||||
)
|
||||
out: Final = await _drive_stream(UnifiedLLMGuardrails(), _ToolRedactingGuardrail(), chunks)
|
||||
out: Final = await _drive_stream(
|
||||
UnifiedLLMGuardrails(), _ToolRedactingGuardrail(require_tool_context=require_tool_context), chunks
|
||||
)
|
||||
calls: Final = tuple(
|
||||
call for chunk in out for choice in chunk.choices for call in choice.delta.tool_calls or ()
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue