mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(guardrails): refresh shared response context for each choice
This commit is contained in:
parent
74608c6859
commit
f4480c903f
2 changed files with 88 additions and 19 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue