mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(guardrails): finish stream checks before releasing tool calls
This commit is contained in:
parent
e80c20aab7
commit
080b5a486a
2 changed files with 34 additions and 6 deletions
|
|
@ -765,7 +765,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
if saw_text_content:
|
||||
async for out in _round(last_chunk, is_final=False):
|
||||
yield out
|
||||
for tool_only in (
|
||||
tool_chunks: Final = tuple(
|
||||
self._tool_call_passthrough_chunk(
|
||||
buffered_item,
|
||||
finish_reason_per_choice=finish_reason_per_choice,
|
||||
|
|
@ -773,9 +773,31 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
)
|
||||
for buffered_item in responses_so_far
|
||||
if self._chunk_has_tool_calls(buffered_item)
|
||||
):
|
||||
)
|
||||
|
||||
async def checked_tail() -> AsyncGenerator[object, None]:
|
||||
try:
|
||||
async for tail_chunk in self._emit_stream_tail(
|
||||
last_chunk=last_chunk,
|
||||
final_round=_round,
|
||||
responses_so_far=responses_so_far,
|
||||
responses_yielded=responses_yielded,
|
||||
):
|
||||
yield tail_chunk
|
||||
except _StreamTerminated as exc:
|
||||
yield exc
|
||||
|
||||
tail_chunks: Final = tuple([chunk async for chunk in checked_tail()])
|
||||
if tail_chunks and isinstance(tail_chunks[-1], _StreamTerminated):
|
||||
for error_chunk in tail_chunks[:-1]:
|
||||
yield error_chunk
|
||||
return
|
||||
for tool_only in tool_chunks:
|
||||
responses_yielded.append(tool_only)
|
||||
yield tool_only
|
||||
for tail_chunk in tail_chunks:
|
||||
yield tail_chunk
|
||||
return
|
||||
|
||||
async for out in self._emit_stream_tail(
|
||||
last_chunk=last_chunk,
|
||||
|
|
|
|||
|
|
@ -1375,14 +1375,15 @@ class TestStreamingTransform:
|
|||
assert out[-2].choices[0].finish_reason == "stop"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_call_blocking_guardrail_is_enforced(self):
|
||||
@pytest.mark.parametrize(("content", "allowed_scans"), [(None, 0), ("proposal", 0), ("proposal", 1)])
|
||||
async def test_tool_call_blocking_guardrail_is_enforced(self, content: str | None, allowed_scans: int):
|
||||
"""A guardrail that blocks on tool calls must terminate the incremental_diff
|
||||
stream: tool calls go through the block decision, not bypass it."""
|
||||
from litellm.exceptions import GuardrailRaisedException
|
||||
|
||||
class _ToolCallBlocker(_StreamingTextGuardrail):
|
||||
async def apply_guardrail(self, inputs, request_data, input_type, **kwargs):
|
||||
if input_type == "response" and inputs.get("tool_calls"):
|
||||
if input_type == "response" and self.response_calls >= allowed_scans:
|
||||
raise GuardrailRaisedException(
|
||||
guardrail_name="tc-block",
|
||||
message="blocked tool call",
|
||||
|
|
@ -1395,7 +1396,7 @@ class TestStreamingTransform:
|
|||
StreamingChoices(
|
||||
index=0,
|
||||
delta=Delta(
|
||||
content=None,
|
||||
content=content,
|
||||
tool_calls=[
|
||||
{
|
||||
"index": 0,
|
||||
|
|
@ -1418,8 +1419,13 @@ class TestStreamingTransform:
|
|||
response=upstream(),
|
||||
request_data={"guardrail_to_apply": _ToolCallBlocker(), "model": "gpt-4"},
|
||||
)
|
||||
async def consume_checked_stream() -> None:
|
||||
async for chunk in stream:
|
||||
assert isinstance(chunk, ModelResponseStream)
|
||||
assert all(not choice.delta.tool_calls and choice.finish_reason is None for choice in chunk.choices)
|
||||
|
||||
with pytest.raises(GuardrailRaisedException):
|
||||
await anext(stream)
|
||||
await consume_checked_stream()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mixed_content_and_tool_call_chunk_does_not_leak_text(self):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue