mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(guardrails): preserve per-choice tool inspection context
This commit is contained in:
parent
4a709fa924
commit
b2d14f9174
3 changed files with 63 additions and 11 deletions
|
|
@ -657,13 +657,22 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
model_response: Final = self._rebuild_ended_stream_per_choice(responses_so_far, litellm_logging_obj)
|
||||
pre_guardrail_texts: Final = self._string_choice_contents(model_response)
|
||||
pre_guardrail_tool_calls: Final = self._function_tool_call_shapes(model_response)
|
||||
await self.process_output_response(
|
||||
response=model_response,
|
||||
guardrail_to_apply=guardrail_to_apply,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_data=request_data,
|
||||
inspection_responses: Final = (
|
||||
tuple(
|
||||
model_response.model_copy(update=MappingProxyType({"choices": [choice]}))
|
||||
for choice in model_response.choices
|
||||
)
|
||||
if pre_guardrail_tool_calls and len(model_response.choices) > 1
|
||||
else (model_response,)
|
||||
)
|
||||
for inspection_response in inspection_responses:
|
||||
await self.process_output_response(
|
||||
response=inspection_response,
|
||||
guardrail_to_apply=guardrail_to_apply,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_data=request_data,
|
||||
)
|
||||
if not deliver_ended_stream_rewrites:
|
||||
return
|
||||
await self._write_ended_stream_text_rewrites(
|
||||
|
|
@ -798,6 +807,24 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
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)
|
||||
if len(assembled.choices) > 1:
|
||||
choice_sinks: Final = tuple((choice.index, StreamTransformSink()) for choice in assembled.choices)
|
||||
for index, choice_sink in choice_sinks:
|
||||
await self._process_streaming_transform(
|
||||
responses_so_far=[self._narrowed_to_choice(chunk, index) for chunk in responses_so_far],
|
||||
guardrail_to_apply=guardrail_to_apply,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_data=request_data,
|
||||
sink=choice_sink,
|
||||
)
|
||||
sink.mutated_text_per_choice = dict(
|
||||
chain.from_iterable(choice_sink.mutated_text_per_choice.items() for _, choice_sink in choice_sinks)
|
||||
)
|
||||
sink.holdback_per_choice = dict(
|
||||
chain.from_iterable(choice_sink.holdback_per_choice.items() for _, choice_sink in choice_sinks)
|
||||
)
|
||||
return
|
||||
tool_calls: Final = chain.from_iterable(choice.message.tool_calls or () for choice in assembled.choices)
|
||||
inputs["tool_calls"] = TypeAdapter(list[ChatCompletionToolCallChunk]).validate_python(
|
||||
tuple(
|
||||
|
|
|
|||
|
|
@ -962,9 +962,11 @@ class _ToolRedactingGuardrail(CustomGuardrail):
|
|||
calls: Final = TypeAdapter(tuple[ChatCompletionMessageToolCall, ...]).validate_python(
|
||||
inputs.get("tool_calls", ())
|
||||
)
|
||||
input_texts: Final = 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", ())
|
||||
"checked:" + text.replace("SECRET", "MASKED")
|
||||
if not self.require_tool_context or (calls and index == len(input_texts) - 1) else text
|
||||
for index, text in enumerate(input_texts)
|
||||
)
|
||||
return {
|
||||
**inputs,
|
||||
|
|
@ -1021,6 +1023,29 @@ class TestStreamingTransform:
|
|||
assert out[-1].usage.total_tokens == 18
|
||||
assert all("SECRET" not in chunk.model_dump_json() for chunk in out)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("tool_choice_index", [0, 1])
|
||||
async def test_tool_dependent_text_rewrites_keep_choice_context(self, tool_choice_index: int) -> None:
|
||||
chunks: Final = (
|
||||
ModelResponseStream(choices=[StreamingChoices(
|
||||
index=index, delta=Delta(content="SECRET" if index == tool_choice_index else "plain"),
|
||||
) for index in (1, 0)]),
|
||||
ModelResponseStream(choices=[StreamingChoices(
|
||||
index=tool_choice_index,
|
||||
delta=Delta(tool_calls=[{
|
||||
"index": 0, "id": "call_context", "type": "function",
|
||||
"function": {"name": "contact", "arguments": '{"contact":"SECRET"}'},
|
||||
}]), finish_reason="tool_calls",
|
||||
)]),
|
||||
)
|
||||
out: Final = await _drive_stream(UnifiedLLMGuardrails(), _ToolRedactingGuardrail(True), chunks)
|
||||
for index in (0, 1):
|
||||
text: Final = "".join(
|
||||
choice.delta.content or "" for chunk in out for choice in chunk.choices if choice.index == index
|
||||
)
|
||||
assert text == ("checked:MASKED" if index == tool_choice_index else "plain")
|
||||
assert all("SECRET" not in chunk.model_dump_json() for chunk in out)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_rewrites_keep_completion_choices_separate(self) -> None:
|
||||
async def response() -> AsyncIterator[ModelResponseStream]:
|
||||
|
|
|
|||
|
|
@ -1395,9 +1395,9 @@ class TestOpenAIChatCompletionsHandlerStreamingOutput:
|
|||
)
|
||||
|
||||
assert [
|
||||
(tool_call["id"], tool_call["function"]["arguments"])
|
||||
for tool_call in guardrail.seen_inputs[-1]["tool_calls"]
|
||||
] == [("call_1", '{"fruit": "persimmon"}'), ("call_2", '{"fruit": "durian"}')]
|
||||
[(tool_call["id"], tool_call["function"]["arguments"]) for tool_call in inputs["tool_calls"]]
|
||||
for inputs in guardrail.seen_inputs
|
||||
] == [[("call_1", '{"fruit": "persimmon"}')], [("call_2", '{"fruit": "durian"}')]]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deliver_ended_stream_tool_rewrites_keep_choice_indices(self) -> None:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue