fix(guardrails): refresh shared response context for each choice

This commit is contained in:
albertbausili 2026-09-28 13:20:57 +02:00
parent 74608c6859
commit f4480c903f
2 changed files with 88 additions and 19 deletions

View file

@ -665,14 +665,28 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
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,
)
inspection_request_data: Final = request_data if request_data is not None else {}
try:
for inspection_response in inspection_responses:
inspection_request_data["response"] = inspection_response
inspection_request_data["responses"] = (
[
self._narrowed_to_choice(chunk, inspection_response.choices[0].index)
for chunk in responses_so_far
]
if len(inspection_response.choices) == 1
else responses_so_far
)
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=inspection_request_data,
)
finally:
inspection_request_data["response"] = model_response
inspection_request_data["responses"] = responses_so_far
if not deliver_ended_stream_rewrites:
return
await self._write_ended_stream_text_rewrites(
@ -807,22 +821,38 @@ 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)
request_data["response"] = assembled
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,
choice_rounds: Final = tuple(
(
choice,
StreamTransformSink(),
[self._narrowed_to_choice(chunk, choice.index) for chunk in responses_so_far],
)
for choice in assembled.choices
)
try:
for choice, choice_sink, choice_chunks in choice_rounds:
request_data["response"] = assembled.model_copy(update=MappingProxyType({"choices": [choice]}))
request_data["responses"] = choice_chunks
await self._process_streaming_transform(
responses_so_far=choice_chunks,
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,
)
finally:
request_data["response"] = assembled
request_data["responses"] = responses_so_far
sink.mutated_text_per_choice = dict(
chain.from_iterable(choice_sink.mutated_text_per_choice.items() for _, choice_sink in choice_sinks)
chain.from_iterable(
choice_sink.mutated_text_per_choice.items() for _, choice_sink, _ in choice_rounds
)
)
sink.holdback_per_choice = dict(
chain.from_iterable(choice_sink.holdback_per_choice.items() for _, choice_sink in choice_sinks)
chain.from_iterable(choice_sink.holdback_per_choice.items() for _, choice_sink, _ in choice_rounds)
)
return
tool_calls: Final = chain.from_iterable(choice.message.tool_calls or () for choice in assembled.choices)

View file

@ -1399,6 +1399,45 @@ class TestOpenAIChatCompletionsHandlerStreamingOutput:
for inputs in guardrail.seen_inputs
] == [[("call_1", '{"fruit": "persimmon"}')], [("call_2", '{"fruit": "durian"}')]]
@pytest.mark.asyncio
@pytest.mark.parametrize("transform", [False, True])
async def test_each_choice_refreshes_shared_guardrail_response_context(self, transform: bool) -> None:
from fastapi import HTTPException
from litellm.llms.base_llm.guardrail_translation.base_translation import StreamTransformSink
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices
class ContextGuardrail(CustomGuardrail):
async def apply_guardrail(
self, inputs: GenericGuardrailAPIInputs, request_data: dict[str, object],
input_type: Literal["request", "response"], logging_obj: object = None,
) -> GenericGuardrailAPIInputs:
response: Final = request_data["response"]
assert isinstance(response, ModelResponse)
request_data["visited_choices"] = (*request_data.get("visited_choices", ()), response.choices[0].index)
if response.choices[0].message.content == "forbidden":
raise HTTPException(status_code=400, detail="later choice blocked")
return inputs
chunks: Final = [ModelResponseStream(choices=[StreamingChoices(
index=index, delta=Delta(content=text, tool_calls=[{
"index": 0, "id": f"call_{index}", "type": "function",
"function": {"name": "contact", "arguments": "{}"},
}]), finish_reason="tool_calls",
) for index, text in enumerate(("allowed", "forbidden"))])]
request_data: Final = {"metadata": {"trace": "retained"}}
with pytest.raises(HTTPException, match="later choice blocked"):
await OpenAIChatCompletionsHandler().process_output_streaming_response(
responses_so_far=chunks, guardrail_to_apply=ContextGuardrail(guardrail_name="context"),
request_data=request_data, deliver_ended_stream_rewrites=True,
stream_transform_sink=StreamTransformSink() if transform else None,
)
assert request_data["visited_choices"] == (0, 1)
assert request_data["metadata"]["trace"] == "retained"
assert request_data["responses"] is chunks
restored: Final = request_data["response"]
assert isinstance(restored, ModelResponse)
assert tuple(choice.index for choice in restored.choices) == (0, 1)
@pytest.mark.asyncio
async def test_deliver_ended_stream_tool_rewrites_keep_choice_indices(self) -> None:
handler: Final = OpenAIChatCompletionsHandler()