From cd42eebbaec048fba8f515523d7e0d046e4aa1a2 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 12 Jun 2026 06:15:48 +0000 Subject: [PATCH] fix(proxy): close guardrail bypass via tool result text and default-mode post-call guardrails on bedrock passthrough Pre-call extraction only read top-level Converse text blocks, so blocked content placed under toolResult.content[].text was forwarded to Bedrock without the key/team guardrail seeing it. Extraction now walks nested tool result text and write-back mutates the owning block in place. Post-call buffering for passthrough used _has_post_call_guardrails, which excludes event_hook=None guardrails. Those guardrails run at post_call, so their output processing was skipped and the raw upstream body was returned. Add a passthrough-specific predicate that counts them. --- .../guardrail_translation/handler.py | 64 +++++++------- litellm/proxy/common_request_processing.py | 21 ++++- .../guardrail_translation/test_handler.py | 86 +++++++++++++++++-- .../test_deferred_guardrail_logging.py | 44 ++++++++++ .../proxy/test_common_request_processing.py | 13 +-- 5 files changed, 177 insertions(+), 51 deletions(-) diff --git a/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py b/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py index f798869764..3a66d39b16 100644 --- a/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py +++ b/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py @@ -26,61 +26,55 @@ def _extract_converse_texts( body: dict, skip_system: bool, skip_tool: bool, -) -> Tuple[List[str], List[Tuple[str, int, int]]]: +) -> Tuple[List[str], List[dict]]: """ Walk a Bedrock Converse request body and collect text content. - Returns (texts, task_mappings) where each task_mapping is - ("system", block_idx, -1) or ("message", msg_idx, content_idx). + Returns (texts, holders) where each holder is the dict that owns the + extracted ``text`` key, so write-back mutates it in place. Tool result + text lives under ``toolResult.content[].text`` rather than the top-level + block and is scanned too unless tool blocks are skipped. """ texts: List[str] = [] - task_mappings: List[Tuple[str, int, int]] = [] + holders: List[dict] = [] + + def _collect(block: dict) -> None: + text = block.get("text") + if text: + texts.append(text) + holders.append(block) if not skip_system: - for i, block in enumerate(body.get("system") or []): - text = block.get("text") if isinstance(block, dict) else None - if text: - texts.append(text) - task_mappings.append(("system", i, -1)) + for block in body.get("system") or []: + if isinstance(block, dict): + _collect(block) - for msg_idx, message in enumerate(body.get("messages") or []): + for message in body.get("messages") or []: if not isinstance(message, dict): continue - for content_idx, block in enumerate(message.get("content") or []): + for block in message.get("content") or []: if not isinstance(block, dict): continue if skip_tool and ("toolUse" in block or "toolResult" in block): continue - text = block.get("text") - if text: - texts.append(text) - task_mappings.append(("message", msg_idx, content_idx)) + _collect(block) + tool_result = block.get("toolResult") + if isinstance(tool_result, dict): + for inner in tool_result.get("content") or []: + if isinstance(inner, dict): + _collect(inner) - return texts, task_mappings + return texts, holders def _write_back_texts( - body: dict, guardrailed_texts: List[str], - task_mappings: List[Tuple[str, int, int]], + holders: List[dict], ) -> None: - for idx, mapping in enumerate(task_mappings): + for idx, holder in enumerate(holders): if idx >= len(guardrailed_texts): break - location, outer_idx, inner_idx = mapping - if location == "system": - system = body.get("system") - if system and isinstance(system, list) and outer_idx < len(system): - system[outer_idx]["text"] = guardrailed_texts[idx] - else: - messages = body.get("messages") - if not ( - messages and isinstance(messages, list) and outer_idx < len(messages) - ): - continue - content = messages[outer_idx].get("content") - if content and isinstance(content, list) and inner_idx < len(content): - content[inner_idx]["text"] = guardrailed_texts[idx] + holder["text"] = guardrailed_texts[idx] class BedrockPassthroughGuardrailHandler(BaseTranslation): @@ -240,7 +234,7 @@ class BedrockPassthroughGuardrailHandler(BaseTranslation): skip_system = effective_skip_system_message_for_guardrail(guardrail_to_apply) skip_tool = effective_skip_tool_message_for_guardrail(guardrail_to_apply) - texts, task_mappings = _extract_converse_texts(body, skip_system, skip_tool) + texts, holders = _extract_converse_texts(body, skip_system, skip_tool) if not texts: return data @@ -259,7 +253,7 @@ class BedrockPassthroughGuardrailHandler(BaseTranslation): guardrailed_texts = guardrailed_inputs.get("texts", []) if guardrailed_texts: - _write_back_texts(body, guardrailed_texts, task_mappings) + _write_back_texts(guardrailed_texts, holders) return data diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index e034dc04ee..d855730f54 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1338,7 +1338,7 @@ class ProxyBaseLLMRequestProcessing: else: generator = response - if self._has_post_call_guardrails(): + if self._has_post_call_guardrails_for_passthrough(): body_bytes = b"".join( [chunk async for chunk in generator] # type: ignore[union-attr] ) @@ -1748,6 +1748,23 @@ class ProxyBaseLLMRequestProcessing: return True return False + @staticmethod + def _has_post_call_guardrails_for_passthrough() -> bool: + """ + True when any guardrail runs at post_call for passthrough responses. + + Unlike _has_post_call_guardrails, an event_hook=None guardrail counts: + should_run_guardrail treats it as matching every hook (post_call + included), so skipping the passthrough buffer here would forward the + raw upstream body and bypass that guardrail's output processing. + """ + for cb in litellm.callbacks: + if not isinstance(cb, CustomGuardrail): + continue + if cb._event_hook_is_event_type(GuardrailEventHooks.post_call): + return True + return False + async def _handle_non_streaming_allm_passthrough_route( self, response: Any, @@ -1755,7 +1772,7 @@ class ProxyBaseLLMRequestProcessing: user_api_key_dict: "UserAPIKeyAuth", custom_headers: dict, ) -> Optional[Response]: - if not self._has_post_call_guardrails(): + if not self._has_post_call_guardrails_for_passthrough(): return None import json as _json diff --git a/tests/test_litellm/llms/bedrock/passthrough/guardrail_translation/test_handler.py b/tests/test_litellm/llms/bedrock/passthrough/guardrail_translation/test_handler.py index 80c29b3f1d..2e613f9099 100644 --- a/tests/test_litellm/llms/bedrock/passthrough/guardrail_translation/test_handler.py +++ b/tests/test_litellm/llms/bedrock/passthrough/guardrail_translation/test_handler.py @@ -72,19 +72,19 @@ class TestExtractConverseTexts: "system": [{"text": "sys text"}], "messages": [{"role": "user", "content": [{"text": "user text"}]}], } - texts, mappings = _extract_converse_texts(body, skip_system=False, skip_tool=False) + texts, holders = _extract_converse_texts(body, skip_system=False, skip_tool=False) assert texts == ["sys text", "user text"] - assert mappings[0] == ("system", 0, -1) - assert mappings[1] == ("message", 0, 0) + assert holders[0] is body["system"][0] + assert holders[1] is body["messages"][0]["content"][0] def test_skip_system(self): body = { "system": [{"text": "sys text"}], "messages": [{"role": "user", "content": [{"text": "user text"}]}], } - texts, mappings = _extract_converse_texts(body, skip_system=True, skip_tool=False) + texts, holders = _extract_converse_texts(body, skip_system=True, skip_tool=False) assert texts == ["user text"] - assert all(m[0] == "message" for m in mappings) + assert holders == [body["messages"][0]["content"][0]] def test_skip_tool_blocks(self): body = { @@ -107,6 +107,30 @@ class TestExtractConverseTexts: texts, _ = _extract_converse_texts(body, skip_system=False, skip_tool=True) assert texts == ["hello"] + def test_extracts_nested_tool_result_text(self): + body = { + "messages": [ + { + "role": "user", + "content": [ + {"text": "hello"}, + { + "toolResult": { + "toolUseId": "1", + "content": [ + {"text": "blocked tool text"}, + {"json": {"k": "v"}}, + ], + } + }, + ], + } + ] + } + texts, holders = _extract_converse_texts(body, skip_system=False, skip_tool=False) + assert texts == ["hello", "blocked tool text"] + assert holders[1] is body["messages"][0]["content"][1]["toolResult"]["content"][0] + def test_non_text_content_blocks_ignored(self): body = { "messages": [ @@ -123,14 +147,36 @@ class TestExtractConverseTexts: class TestWriteBackTexts: def test_writes_system_text(self): body = {"system": [{"text": "original"}], "messages": []} - _write_back_texts(body, ["replaced"], [("system", 0, -1)]) + _, holders = _extract_converse_texts(body, skip_system=False, skip_tool=False) + _write_back_texts(["replaced"], holders) assert body["system"][0]["text"] == "replaced" def test_writes_message_text(self): body = {"messages": [{"role": "user", "content": [{"text": "original"}]}]} - _write_back_texts(body, ["replaced"], [("message", 0, 0)]) + _, holders = _extract_converse_texts(body, skip_system=False, skip_tool=False) + _write_back_texts(["replaced"], holders) assert body["messages"][0]["content"][0]["text"] == "replaced" + def test_writes_nested_tool_result_text(self): + body = { + "messages": [ + { + "role": "user", + "content": [ + { + "toolResult": { + "toolUseId": "1", + "content": [{"text": "original"}], + } + } + ], + } + ] + } + _, holders = _extract_converse_texts(body, skip_system=False, skip_tool=False) + _write_back_texts(["masked"], holders) + assert body["messages"][0]["content"][0]["toolResult"]["content"][0]["text"] == "masked" + def test_extra_non_text_fields_untouched(self): body = { "messages": [ @@ -151,7 +197,8 @@ class TestWriteBackTexts: "inferenceConfig": {"maxTokens": 100}, } original = copy.deepcopy(body) - _write_back_texts(body, ["replaced"], [("message", 0, 0)]) + _, holders = _extract_converse_texts(body, skip_system=False, skip_tool=False) + _write_back_texts(["replaced"], holders) assert body["messages"][0]["content"][0]["text"] == "replaced" assert body["messages"][0]["content"][1] == original["messages"][0]["content"][1] assert body["inferenceConfig"] == original["inferenceConfig"] @@ -203,6 +250,29 @@ class TestBedrockPassthroughGuardrailHandlerInput: with pytest.raises(HTTPException): await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + @pytest.mark.asyncio + async def test_tool_result_text_scanned_and_masked(self): + handler = BedrockPassthroughGuardrailHandler() + data = _converse_data() + data["data"]["messages"][0]["content"].append( + { + "toolResult": { + "toolUseId": "t1", + "content": [{"text": "My SSN is 123-45-6789"}], + } + } + ) + guardrail = _make_guardrail( + {"texts": ["You are helpful.", "Hello world", "[REDACTED]"]} + ) + + result = await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + sent_texts = guardrail.apply_guardrail.call_args.kwargs["inputs"]["texts"] + assert "My SSN is 123-45-6789" in sent_texts + tool_result = result["data"]["messages"][0]["content"][2]["toolResult"] + assert tool_result["content"][0]["text"] == "[REDACTED]" + @pytest.mark.asyncio async def test_non_converse_endpoint_skips_apply_guardrail(self): handler = BedrockPassthroughGuardrailHandler() diff --git a/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py b/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py index e9ff193e04..d104f8d7f4 100644 --- a/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py +++ b/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py @@ -154,6 +154,50 @@ class TestHasPostCallGuardrails: assert ProxyBaseLLMRequestProcessing._has_post_call_guardrails() is False +class TestHasPostCallGuardrailsForPassthrough: + """Passthrough buffering must include event_hook=None guardrails. + + Those guardrails run at post_call (should_run_guardrail treats None as + matching every hook); skipping the buffer would forward the raw upstream + body and bypass output processing. + """ + + def test_returns_true_for_event_hook_none(self): + with patch("litellm.callbacks", [AllEventsGuardrail()]): + assert ( + ProxyBaseLLMRequestProcessing._has_post_call_guardrails_for_passthrough() + is True + ) + + def test_returns_true_for_post_call_guardrail(self): + with patch("litellm.callbacks", [PostCallGuardrail()]): + assert ( + ProxyBaseLLMRequestProcessing._has_post_call_guardrails_for_passthrough() + is True + ) + + def test_returns_false_for_pre_call_only(self): + with patch("litellm.callbacks", [PreCallGuardrail()]): + assert ( + ProxyBaseLLMRequestProcessing._has_post_call_guardrails_for_passthrough() + is False + ) + + def test_returns_false_for_no_callbacks(self): + with patch("litellm.callbacks", []): + assert ( + ProxyBaseLLMRequestProcessing._has_post_call_guardrails_for_passthrough() + is False + ) + + def test_ignores_non_guardrail_callbacks(self): + with patch("litellm.callbacks", ["langfuse", CustomLogger()]): + assert ( + ProxyBaseLLMRequestProcessing._has_post_call_guardrails_for_passthrough() + is False + ) + + # --------------------------------------------------------------------------- # 2. Non-streaming: deferral flag → closure stored, create_task skipped # --------------------------------------------------------------------------- diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index b2d9ed8507..2c9984fe21 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -2394,7 +2394,7 @@ class TestAllmPassthroughRoutePostCallGuardrails: proxy_logging_obj = ProxyLogging(user_api_key_cache=MagicMock()) monkeypatch.setattr(proxy_logging_obj, "post_call_success_hook", capture_hook) - with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails", return_value=True): + with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True): processing_obj = ProxyBaseLLMRequestProcessing(data={}) result = await processing_obj._handle_non_streaming_allm_passthrough_route( response=httpx_response, @@ -2443,7 +2443,7 @@ class TestAllmPassthroughRoutePostCallGuardrails: proxy_logging_obj = ProxyLogging(user_api_key_cache=MagicMock()) monkeypatch.setattr(proxy_logging_obj, "post_call_success_hook", non_dict_hook) - with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails", return_value=True): + with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True): processing_obj = ProxyBaseLLMRequestProcessing(data={}) result = await processing_obj._handle_non_streaming_allm_passthrough_route( response=httpx_response, @@ -2499,8 +2499,9 @@ class TestAllmPassthroughRoutePostCallGuardrails: @pytest.mark.asyncio async def test_no_aread_when_no_post_call_guardrails(self, monkeypatch): """ - When _has_post_call_guardrails() is False the httpx response must not be - read — the caller handles streaming or error paths normally. + When _has_post_call_guardrails_for_passthrough() is False the httpx + response must not be read — the caller handles streaming or error paths + normally. """ import json @@ -2519,7 +2520,7 @@ class TestAllmPassthroughRoutePostCallGuardrails: hook_spy = AsyncMock() monkeypatch.setattr(proxy_logging_obj, "post_call_success_hook", hook_spy) - with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails", return_value=False): + with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=False): processing_obj = ProxyBaseLLMRequestProcessing(data={}) result = await processing_obj._handle_non_streaming_allm_passthrough_route( response=httpx_response, @@ -2629,7 +2630,7 @@ class TestEventStreamAllmPassthroughRoute: "content-length": "99", } - with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails", return_value=True): + with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True): processing_obj = ProxyBaseLLMRequestProcessing(data={}) result = await processing_obj._handle_non_streaming_allm_passthrough_route( response=mock_response,