fix(tool_permission): fail closed on unverifiable SSE streams and end the turn when every tool call is denied

An SSE stream that cannot be positively identified as Anthropic (no
parseable message_start event) now blocks instead of passing through
unscanned, closing the bypass where any raw-SSE backend skipped tool
permission checks entirely. Buffered chunks are joined back into one
stream before parsing, so events split across network chunk boundaries
assemble correctly instead of being silently dropped. Rewrite mode now
resets finish_reason to stop when no tool call survives, so the
re-encoded Anthropic stream reports stop_reason end_turn and clients do
not wait for a tool result that never comes
This commit is contained in:
mateo-berri 2026-08-05 15:06:51 -07:00
parent bee787b4b5
commit f16f3e23cd
2 changed files with 107 additions and 3 deletions

View file

@ -734,6 +734,13 @@ class ToolPermissionGuardrail(CustomGuardrail):
else:
choice.message.content = "\n".join(error_messages)
if (
not choice.message.tool_calls
and getattr(choice.message, "function_call", None) is None
and choice.finish_reason in ("tool_calls", "function_call")
):
choice.finish_reason = "stop"
@log_guardrail_information
async def async_pre_call_hook(
self,
@ -878,6 +885,14 @@ class ToolPermissionGuardrail(CustomGuardrail):
anthropic_response: Final = self._assemble_anthropic_stream(all_chunks)
if anthropic_response is None:
if self._is_raw_sse_stream(all_chunks):
raise GuardrailRaisedException(
guardrail_name=self.guardrail_name,
message=(
"Streamed response could not be verified for tool permissions "
"(not a parseable Anthropic SSE stream), blocking it"
),
)
for chunk in all_chunks:
yield chunk
return
@ -910,18 +925,42 @@ class ToolPermissionGuardrail(CustomGuardrail):
verbose_proxy_logger.debug("Tool Permission Guardrail Post-Call Hook: All tools allowed")
return denied_tools
@staticmethod
def _joined_sse_stream(all_chunks: Sequence[Any]) -> str | None:
raw: Final = b"".join(
chunk if isinstance(chunk, bytes) else chunk.encode("utf-8")
for chunk in all_chunks
if isinstance(chunk, (str, bytes))
)
try:
return raw.decode("utf-8")
except UnicodeDecodeError:
return None
@staticmethod
def _has_anthropic_message_start(sse_stream: str) -> bool:
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import (
AnthropicPassthroughLoggingHandler,
)
return any(
(event_data := AnthropicPassthroughLoggingHandler._extract_sse_data(event)) is not None # pyright: ignore[reportPrivateUsage] # same parser the assembler uses; a private import beats forking SSE parsing
and event_data.get("type") == "message_start"
for event in AnthropicPassthroughLoggingHandler._split_sse_chunk_into_events(sse_stream) # pyright: ignore[reportPrivateUsage] # same parser the assembler uses
)
@staticmethod
def _assemble_anthropic_stream(all_chunks: Sequence[Any]) -> ModelResponse | None:
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import (
AnthropicPassthroughLoggingHandler,
)
sse_chunks: Final = tuple(chunk for chunk in all_chunks if isinstance(chunk, (str, bytes)))
if not sse_chunks:
sse_stream: Final = ToolPermissionGuardrail._joined_sse_stream(all_chunks)
if sse_stream is None or not ToolPermissionGuardrail._has_anthropic_message_start(sse_stream):
return None
try:
assembled = AnthropicPassthroughLoggingHandler._build_complete_streaming_response( # pyright: ignore[reportPrivateUsage] # the only SSE-to-ModelResponse assembler; reimplementing it here would fork the parser
all_chunks=sse_chunks,
all_chunks=(sse_stream,),
litellm_logging_obj=None, # pyright: ignore[reportArgumentType] # only forwarded to stream_chunk_builder, which accepts None
model="",
)

View file

@ -752,6 +752,26 @@ class TestToolPermissionGuardrail:
assert isinstance(choice.message.content, str)
assert "Permission denied" in choice.message.content
def test_modify_response_resets_finish_reason_when_every_tool_call_is_denied(self):
tool_call = ChatCompletionMessageToolCall(function={"name": "Read", "arguments": "{}"}, id="call_123")
response = ModelResponse(
choices=[Choices(finish_reason="tool_calls", message={"tool_calls": [tool_call], "content": ""})]
)
denied_tools = [
(
tool_call,
PermissionError(tool_name="Read", rule_id="deny_read", message="Tool 'Read' denied by rule 'deny_read'"),
)
]
self.guardrail._modify_response_with_permission_errors(response, denied_tools)
choice = response.choices[0]
assert isinstance(choice, Choices)
assert choice.finish_reason == "stop", (
"keeping finish_reason tool_calls with no surviving tool calls leaves the client waiting on a tool"
)
def test_modify_response_with_permission_errors_filters_legacy_function_call(self):
response = ModelResponse(
choices=[
@ -1190,3 +1210,48 @@ class TestToolPermissionGuardrailAnthropicMessages:
body = b"".join(c if isinstance(c, bytes) else str(c).encode() for c in out).decode()
assert '"type": "tool_use"' not in body, "denied tool_use must not survive into the rewritten stream"
assert "Permission denied" in body
assert '"stop_reason": "end_turn"' in body, (
"dropping every tool_use must end the turn, or the client waits for a tool result that never comes"
)
assert '"stop_reason": "tool_use"' not in body
def _resplit(self, chunks, size=7):
joined = b"".join(chunks)
return [joined[i : i + size] for i in range(0, len(joined), size)]
@pytest.mark.asyncio
async def test_denied_tool_use_is_caught_when_sse_events_are_split_across_chunk_boundaries(self):
with patch.object(self.blocking, "should_run_guardrail", return_value=True):
with pytest.raises(GuardrailRaisedException) as exc_info:
await self._drain(self.blocking, self._resplit(self._sse_chunks("Read")))
assert "deny_read" in str(exc_info.value), (
"a stream split mid-event must still assemble and hit the rule, not fail as unparseable"
)
@pytest.mark.asyncio
async def test_allowed_stream_split_across_chunk_boundaries_is_passed_through_verbatim(self):
chunks = self._resplit(self._sse_chunks("Bash"))
with patch.object(self.blocking, "should_run_guardrail", return_value=True):
out = await self._drain(self.blocking, chunks)
assert out == chunks
@pytest.mark.asyncio
async def test_non_anthropic_sse_stream_fails_closed(self):
gemini_chunks = [
b'data: {"candidates": [{"content": {"parts": [{"functionCall": '
b'{"name": "run_shell", "args": {"command": "ls"}}}], "role": "model"}}]}\n\n',
b'data: {"candidates": [{"content": {"parts": [{"text": "done"}]}, "finishReason": "STOP"}]}\n\n',
]
with patch.object(self.blocking, "should_run_guardrail", return_value=True):
with pytest.raises(GuardrailRaisedException):
await self._drain(self.blocking, gemini_chunks)
@pytest.mark.asyncio
async def test_unparseable_sse_stream_fails_closed(self):
with patch.object(self.blocking, "should_run_guardrail", return_value=True):
with pytest.raises(GuardrailRaisedException):
await self._drain(self.blocking, [b"data: not-json\n\n", b"event: weird\n\n"])