mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
fix: discard a streamed rewrite that drops or adds a tool call
A guardrail that removes or adds a tool call on an ended stream used to be silently ignored: every handler substitutes the original list on a count mismatch and the executor skipped its observer once the translation could deliver rewrites. The executor now tracks the count change on the observer and releases the original chunks with the discard warning on every translation, matching what the merge base did for any tool call rewrite
This commit is contained in:
parent
1133507565
commit
ddedb4867b
3 changed files with 108 additions and 6 deletions
|
|
@ -78,6 +78,10 @@ def _rewrote(sent: tuple[object, ...] | None, returned: tuple[object, ...] | Non
|
|||
return sent is not None and returned is not None and returned != sent
|
||||
|
||||
|
||||
def _changed_count(sent: tuple[object, ...] | None, returned: tuple[object, ...] | None) -> bool:
|
||||
return sent is not None and returned is not None and len(returned) != len(sent)
|
||||
|
||||
|
||||
_GuardrailMethodT = TypeVar("_GuardrailMethodT", bound=Callable[..., object])
|
||||
|
||||
|
||||
|
|
@ -91,8 +95,9 @@ class _StreamRewriteObserver(CustomGuardrail):
|
|||
guardrail. It records whether the guardrail returned different output than it was given,
|
||||
which for guardrails like Bedrock's ANONYMIZED action is only known at runtime. Text and
|
||||
tool-call rewrites are deliverable on translations that write them back across the
|
||||
buffered chunks (``delivers_ended_stream_rewrites``); rewrites on any other translation
|
||||
are discarded by the executor, which releases the original chunks.
|
||||
buffered chunks (``delivers_ended_stream_rewrites``); rewrites on any other translation,
|
||||
and a rewrite that drops or adds a tool call on any translation, are discarded by the
|
||||
executor, which releases the original chunks.
|
||||
The inner guardrail's ``apply_guardrail`` already records the guardrail information
|
||||
and span, so the observer's stays out of ``log_guardrail_information``."""
|
||||
|
||||
|
|
@ -101,6 +106,7 @@ class _StreamRewriteObserver(CustomGuardrail):
|
|||
self.inner: Final = inner
|
||||
self.rewrote_texts = False
|
||||
self.rewrote_tool_calls = False
|
||||
self.changed_tool_call_count = False
|
||||
|
||||
def structured_messages_cover_full_request(self) -> bool:
|
||||
return self.inner.structured_messages_cover_full_request()
|
||||
|
|
@ -118,9 +124,11 @@ class _StreamRewriteObserver(CustomGuardrail):
|
|||
outputs: Final = await self.inner.apply_guardrail(
|
||||
inputs=inputs, request_data=request_data, input_type=input_type, logging_obj=logging_obj
|
||||
)
|
||||
returned_tool_shapes: Final = _tool_call_shapes(outputs.get("tool_calls"))
|
||||
self.rewrote_texts = self.rewrote_texts or _rewrote(sent_texts, _text_snapshot(outputs.get("texts")))
|
||||
self.rewrote_tool_calls = self.rewrote_tool_calls or _rewrote(
|
||||
sent_tool_shapes, _tool_call_shapes(outputs.get("tool_calls"))
|
||||
self.rewrote_tool_calls = self.rewrote_tool_calls or _rewrote(sent_tool_shapes, returned_tool_shapes)
|
||||
self.changed_tool_call_count = self.changed_tool_call_count or _changed_count(
|
||||
sent_tool_shapes, returned_tool_shapes
|
||||
)
|
||||
return outputs
|
||||
|
||||
|
|
@ -325,7 +333,9 @@ class PipelineExecutor:
|
|||
except UndeliverableStreamRewrite:
|
||||
_release_original_chunks(step.guardrail, streaming_chunks, originals)
|
||||
else:
|
||||
if not deliver_rewrites and (observer.rewrote_texts or observer.rewrote_tool_calls):
|
||||
if observer.changed_tool_call_count or (
|
||||
not deliver_rewrites and (observer.rewrote_texts or observer.rewrote_tool_calls)
|
||||
):
|
||||
_release_original_chunks(step.guardrail, streaming_chunks, originals)
|
||||
if not callback.records_own_guardrail_information:
|
||||
add_guardrail_to_applied_guardrails_header(request_data=hook_input, guardrail_name=step.guardrail)
|
||||
|
|
|
|||
|
|
@ -1106,7 +1106,8 @@ class _WritingTranslation:
|
|||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
responses_so_far[0]["text"] = outputs["texts"][0]
|
||||
responses_so_far[0]["tool_call"] = outputs["tool_calls"][0]
|
||||
if len(outputs["tool_calls"]) == 1:
|
||||
responses_so_far[0]["tool_call"] = outputs["tool_calls"][0]
|
||||
return responses_so_far
|
||||
|
||||
|
||||
|
|
@ -1242,6 +1243,26 @@ async def test_streaming_step_delivers_tool_call_rewrite_through_writing_transla
|
|||
assert not any("discarded" in record.getMessage() for record in caplog.records)
|
||||
|
||||
|
||||
class _ToolCallDroppingGuardrail(CustomGuardrail):
|
||||
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):
|
||||
return {**inputs, "texts": ["hello [MASKED]"], "tool_calls": []}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_step_discards_whole_rewrite_when_guardrail_drops_a_tool_call(monkeypatch, caplog):
|
||||
monkeypatch.setattr(litellm, "callbacks", [_ToolCallDroppingGuardrail()])
|
||||
chunks = [_chunk()]
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
result = await _run_streaming_step(_WritingTranslation(), chunks)
|
||||
|
||||
_assert_passed_with_discard_warning(result, caplog)
|
||||
assert chunks == [_chunk()]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_step_discards_tool_call_rewrite_when_translation_lacks_write_back(monkeypatch, caplog):
|
||||
monkeypatch.setattr(litellm, "callbacks", [_TextAndToolCallRewritingGuardrail(rewrite_tool_call=True)])
|
||||
|
|
|
|||
|
|
@ -2322,3 +2322,74 @@ async def test_streaming_iterator_hook_pipeline_delivers_function_call_rewrite_o
|
|||
assert delivered[4]["item"]["arguments"] == '{"fruit": "[MASKED]"}'
|
||||
assert delivered[5]["response"]["output"][0]["arguments"] == '{"fruit": "[MASKED]"}'
|
||||
assert "persimmon" not in json.dumps(delivered)
|
||||
|
||||
|
||||
def _drop_tool_calls(inputs: Dict[str, Any]) -> Dict[str, Any]:
|
||||
return {"tool_calls": []}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_iterator_hook_pipeline_discards_dropped_tool_call_on_chat_chunks(
|
||||
proxy_logging, make_user_api_key_auth, monkeypatch, caplog
|
||||
):
|
||||
monkeypatch.setattr(litellm, "callbacks", [_rewriting_stream_guardrail(_drop_tool_calls)])
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False)
|
||||
data = _post_call_pipeline_data(stream=True)
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
delivered = [
|
||||
item
|
||||
async for item in proxy_logging.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=make_user_api_key_auth(request_route="/v1/chat/completions"),
|
||||
response=_async_chunk_iter(_tool_call_stream_chunks()),
|
||||
request_data=data,
|
||||
)
|
||||
]
|
||||
|
||||
assert delivered[0].choices[0].delta.tool_calls[0].function.arguments == '{"ssn": "123"}'
|
||||
assert delivered[1].choices[0].finish_reason == "tool_calls"
|
||||
assert any("'gr-post'" in message and "discarded" in message for message in _warnings(caplog))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_iterator_hook_pipeline_discards_dropped_tool_call_on_anthropic_sse(
|
||||
proxy_logging, make_user_api_key_auth, monkeypatch, caplog
|
||||
):
|
||||
monkeypatch.setattr(litellm, "callbacks", [_rewriting_stream_guardrail(_drop_tool_calls)])
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False)
|
||||
data = _post_call_pipeline_data(stream=True)
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
delivered = [
|
||||
item
|
||||
async for item in proxy_logging.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=make_user_api_key_auth(request_route="/v1/messages"),
|
||||
response=_async_chunk_iter(_anthropic_tool_use_sse_chunks()),
|
||||
request_data=data,
|
||||
)
|
||||
]
|
||||
|
||||
assert delivered == _anthropic_tool_use_sse_chunks()
|
||||
assert any("'gr-post'" in message and "discarded" in message for message in _warnings(caplog))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_iterator_hook_pipeline_discards_dropped_tool_call_on_responses_events(
|
||||
proxy_logging, make_user_api_key_auth, monkeypatch, caplog
|
||||
):
|
||||
monkeypatch.setattr(litellm, "callbacks", [_rewriting_stream_guardrail(_drop_tool_calls)])
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False)
|
||||
data = _post_call_pipeline_data(stream=True)
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
delivered = [
|
||||
item
|
||||
async for item in proxy_logging.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=make_user_api_key_auth(request_route="/v1/responses"),
|
||||
response=_async_chunk_iter(_responses_function_call_events()),
|
||||
request_data=data,
|
||||
)
|
||||
]
|
||||
|
||||
assert delivered == _responses_function_call_events()
|
||||
assert any("'gr-post'" in message and "discarded" in message for message in _warnings(caplog))
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue