From 721babb0e789c4a0b781957eca453ca921115eac Mon Sep 17 00:00:00 2001 From: STiFLeR7 Date: Thu, 30 Jul 2026 11:00:35 +0530 Subject: [PATCH] fix(a2a): close guardrail bypass for kind:data parts in A2A protocol extract_text_from_a2a_message folds kind:data parts into the completion text callers see, but A2AGuardrailHandler's input/output extraction only scanned kind:text parts, letting structured data content reach callers without ever being scanned or redacted by output/input guardrails. Extend both process_input_messages and process_output_response (plus the streaming path) to serialize and scan data parts the same way, using a shared serialize_a2a_data_part helper so the completion-text builder and the guardrail extractor can't diverge again. Guardrailed values are written back into the correct field (text or data) per part. Fixes the security finding flagged on this PR: A2A data parts bypass output guardrails (litellm/llms/a2a/common_utils.py:89). --- .../a2a/chat/guardrail_translation/handler.py | 66 ++++++--- litellm/llms/a2a/common_utils.py | 18 ++- tests/test_litellm/llms/a2a/chat/__init__.py | 0 .../chat/guardrail_translation/__init__.py | 0 .../guardrail_translation/test_handler.py | 133 ++++++++++++++++++ 5 files changed, 191 insertions(+), 26 deletions(-) create mode 100644 tests/test_litellm/llms/a2a/chat/__init__.py create mode 100644 tests/test_litellm/llms/a2a/chat/guardrail_translation/__init__.py create mode 100644 tests/test_litellm/llms/a2a/chat/guardrail_translation/test_handler.py diff --git a/litellm/llms/a2a/chat/guardrail_translation/handler.py b/litellm/llms/a2a/chat/guardrail_translation/handler.py index 5c30ff4747a..54abcd9b511 100644 --- a/litellm/llms/a2a/chat/guardrail_translation/handler.py +++ b/litellm/llms/a2a/chat/guardrail_translation/handler.py @@ -17,6 +17,7 @@ from typing import TYPE_CHECKING, Any, Final, Optional from typing_extensions import ReadOnly, TypedDict from litellm._logging import verbose_proxy_logger +from litellm.llms.a2a.common_utils import serialize_a2a_data_part from litellm.llms.base_llm.guardrail_translation.base_translation import ( BaseTranslation, StreamingScanKey, @@ -34,6 +35,7 @@ class _A2ATextPart(TypedDict, total=False): kind: ReadOnly[str] text: ReadOnly[str] + data: ReadOnly[object] class A2AGuardrailHandler(BaseTranslation): @@ -45,8 +47,10 @@ class A2AGuardrailHandler(BaseTranslation): 2. Process output responses (post-call hook) - extracts text from A2A response parts A2A Message Format: - - Input: params.message.parts[].text (where kind == "text") - - Output: result.message.parts[].text or result.artifacts[].parts[].text + - Input: params.message.parts[].text (where kind == "text") or + params.message.parts[].data (where kind == "data") + - Output: result.message.parts[].text or result.artifacts[].parts[].text, + and the "data" equivalents of both """ async def process_input_messages( @@ -78,15 +82,23 @@ class A2AGuardrailHandler(BaseTranslation): return data texts_to_check: Final[list[str]] = [] - text_part_indices: Final[list[int]] = [] # Track which parts contain text + # Track which parts contain scannable content, and which field to write + # the guardrailed value back to ("text" or "data") + part_mappings: Final[list[tuple[int, str]]] = [] - # Step 1: Extract text from all text parts + # Step 1: Extract text from all text parts, and serialized data from all data parts for part_idx, part in enumerate(parts): - if part.get("kind") == "text": + kind = part.get("kind") + if kind == "text": text = part.get("text", "") if text: texts_to_check.append(text) - text_part_indices.append(part_idx) + part_mappings.append((part_idx, "text")) + elif kind == "data": + part_data = part.get("data") + if part_data is not None: + texts_to_check.append(serialize_a2a_data_part(part_data)) + part_mappings.append((part_idx, "data")) # Step 2: Apply guardrail to all texts in batch if texts_to_check: @@ -110,9 +122,9 @@ class A2AGuardrailHandler(BaseTranslation): guardrailed_texts: Final = guardrailed_inputs.get("texts", []) # Step 3: Apply guardrailed text back to original parts - if guardrailed_texts and len(guardrailed_texts) == len(text_part_indices): - for task_idx, part_idx in enumerate(text_part_indices): - parts[part_idx]["text"] = guardrailed_texts[task_idx] + if guardrailed_texts and len(guardrailed_texts) == len(part_mappings): + for task_idx, (part_idx, field) in enumerate(part_mappings): + parts[part_idx][field] = guardrailed_texts[task_idx] verbose_proxy_logger.debug("A2A: Processed input message: %s", message) @@ -164,7 +176,7 @@ class A2AGuardrailHandler(BaseTranslation): texts_to_check: Final[list[str]] = [] # Each mapping is (path_to_parts_list, part_index) # path_to_parts_list is a tuple of keys to navigate to the parts list - task_mappings: Final[list[tuple[tuple[str, ...], int]]] = [] + task_mappings: Final[list[tuple[tuple[str, ...], int, str]]] = [] # Extract texts from all possible locations self._extract_texts_from_result( @@ -205,11 +217,12 @@ class A2AGuardrailHandler(BaseTranslation): # Step 3: Apply guardrailed text back to original response if guardrailed_texts and len(guardrailed_texts) == len(task_mappings): - for task_idx, (path, part_idx) in enumerate(task_mappings): + for task_idx, (path, part_idx, field) in enumerate(task_mappings): self._apply_text_to_path( result=result, path=path, part_idx=part_idx, + field=field, text=guardrailed_texts[task_idx], ) @@ -281,8 +294,8 @@ class A2AGuardrailHandler(BaseTranslation): result = obj.get("result", {}) if not isinstance(result, dict): continue - texts_in_chunk: list[str] = [] - mappings: list[tuple[tuple[str, ...], int]] = [] + texts_in_chunk: Final[list[str]] = [] + mappings: Final[list[tuple[tuple[str, ...], int, str]]] = [] self._extract_texts_from_result( result=result, texts_to_check=texts_in_chunk, @@ -292,20 +305,22 @@ class A2AGuardrailHandler(BaseTranslation): continue if orig_i == first_chunk_with_text: # Put full guardrailed text in first text part; clear others - for task_idx, (path, part_idx) in enumerate(mappings): + for task_idx, (path, part_idx, field) in enumerate(mappings): text = guardrailed_text if task_idx == 0 else "" self._apply_text_to_path( result=result, path=path, part_idx=part_idx, + field=field, text=text, ) else: - for path, part_idx in mappings: + for path, part_idx, field in mappings: self._apply_text_to_path( result=result, path=path, part_idx=part_idx, + field=field, text="", ) @@ -363,7 +378,7 @@ class A2AGuardrailHandler(BaseTranslation): self, result: dict[str, Any], texts_to_check: list[str], - task_mappings: list[tuple[tuple[str, ...], int]], + task_mappings: list[tuple[tuple[str, ...], int, str]], ) -> None: """ Extract text from all possible locations in an A2A result. @@ -433,21 +448,28 @@ class A2AGuardrailHandler(BaseTranslation): parts: Sequence[_A2ATextPart], path: tuple[str, ...], texts_to_check: list[str], - task_mappings: list[tuple[tuple[str, ...], int]], + task_mappings: list[tuple[tuple[str, ...], int, str]], ) -> None: - """Extract text from message parts.""" + """Extract text from message parts, and serialized data from data parts.""" for part_idx, part in enumerate(parts): - if part.get("kind") == "text": + kind = part.get("kind") + if kind == "text": text = part.get("text", "") if text: texts_to_check.append(text) - task_mappings.append((path, part_idx)) + task_mappings.append((path, part_idx, "text")) + elif kind == "data": + data = part.get("data") + if data is not None: + texts_to_check.append(serialize_a2a_data_part(data)) + task_mappings.append((path, part_idx, "data")) def _apply_text_to_path( self, result: dict[str | int, Any], path: tuple[str, ...], part_idx: int, + field: str, text: str, ) -> None: """Apply guardrailed text back to the specified path in the result.""" @@ -460,5 +482,5 @@ class A2AGuardrailHandler(BaseTranslation): else: current = current[key] - # Update the text in the part - current[part_idx]["text"] = text + # Update the guardrailed value in the part + current[part_idx][field] = text diff --git a/litellm/llms/a2a/common_utils.py b/litellm/llms/a2a/common_utils.py index 7dddfb011a0..cc01fde89db 100644 --- a/litellm/llms/a2a/common_utils.py +++ b/litellm/llms/a2a/common_utils.py @@ -63,6 +63,19 @@ def convert_messages_to_prompt(messages: list[AllMessageValues]) -> str: return "\n".join(conversation_parts) +def serialize_a2a_data_part(data: Any) -> str: + """ + Serialize an A2A ``data``-kind part's payload to text. + + Used both to build the flattened completion text shown to callers and to + extract guardrail-scannable text, so the two stay in sync. + """ + try: + return json.dumps(data, ensure_ascii=False) + except (TypeError, ValueError): + return str(data) + + def extract_text_from_a2a_message(message: dict[str, Any], depth: int = 0, max_depth: int = 10) -> str: """ Extract text content from A2A message parts. @@ -88,10 +101,7 @@ def extract_text_from_a2a_message(message: dict[str, Any], depth: int = 0, max_d elif kind == "data": data = part.get("data") if data is not None: - try: - text_parts.append(json.dumps(data, ensure_ascii=False)) - except (TypeError, ValueError): - text_parts.append(str(data)) + text_parts.append(serialize_a2a_data_part(data)) # Handle nested parts if they exist elif "parts" in part: nested_text = extract_text_from_a2a_message(part, depth + 1, max_depth) diff --git a/tests/test_litellm/llms/a2a/chat/__init__.py b/tests/test_litellm/llms/a2a/chat/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/a2a/chat/guardrail_translation/__init__.py b/tests/test_litellm/llms/a2a/chat/guardrail_translation/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/a2a/chat/guardrail_translation/test_handler.py b/tests/test_litellm/llms/a2a/chat/guardrail_translation/test_handler.py new file mode 100644 index 00000000000..4f958ba15fd --- /dev/null +++ b/tests/test_litellm/llms/a2a/chat/guardrail_translation/test_handler.py @@ -0,0 +1,133 @@ +""" +Unit tests for A2A Protocol Guardrail Translation Handler + +Regression coverage for the "data"-kind part guardrail bypass: A2A responses +can carry structured content in `kind: "data"` parts, which +`extract_text_from_a2a_message` (used to build the completion text callers +see) folds into the final text, but the guardrail handler previously only +inspected `kind: "text"` parts, so guarded output checks were skipped for +that content path. +""" + +import os +import sys +from typing import Any, Literal, Optional + +import pytest + +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../../.."))) + +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.llms.a2a.chat.guardrail_translation.handler import A2AGuardrailHandler +from litellm.types.utils import GenericGuardrailAPIInputs + + +class MockGuardrail(CustomGuardrail): + """Mock guardrail that uppercases text so we can assert exactly what was scanned and where the result landed.""" + + def __init__(self, guardrail_name: str = "test"): + super().__init__(guardrail_name=guardrail_name) + self.last_inputs: Optional[GenericGuardrailAPIInputs] = None + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional[Any] = None, + ) -> GenericGuardrailAPIInputs: + self.last_inputs = inputs + texts = inputs.get("texts", []) + return {"texts": [text.upper() for text in texts]} + + +@pytest.mark.asyncio +async def test_process_output_response_scans_data_parts(): + """A `kind: data` part in the output must be sent to the guardrail and the + guardrailed value written back into `data`, not silently skipped.""" + handler = A2AGuardrailHandler() + guardrail = MockGuardrail() + + response = { + "result": { + "kind": "message", + "parts": [ + {"kind": "text", "text": "hello"}, + {"kind": "data", "data": {"secret": "leak-me"}}, + ], + } + } + + result = await handler.process_output_response( + response=response, + guardrail_to_apply=guardrail, + ) + + # The data part's serialized content must have reached the guardrail. + assert guardrail.last_inputs is not None + scanned_texts = guardrail.last_inputs["texts"] + assert any("leak-me" in t for t in scanned_texts) + + # The guardrailed (uppercased) value must be written back into "data", + # and the part must remain a "data" part, not be silently dropped or + # converted into an unguarded pass-through. + data_part = result["result"]["parts"][1] + assert data_part["kind"] == "data" + assert "LEAK-ME" in data_part["data"] + + # The text part must still be guardrailed as before (no regression). + text_part = result["result"]["parts"][0] + assert text_part["text"] == "HELLO" + + +@pytest.mark.asyncio +async def test_process_output_response_data_only_still_scanned(): + """A response with ONLY a data part (no text parts at all) must not be + skipped as "no text content in response".""" + handler = A2AGuardrailHandler() + guardrail = MockGuardrail() + + response = { + "result": { + "kind": "message", + "parts": [{"kind": "data", "data": {"result": {"msg": "pong"}}}], + } + } + + result = await handler.process_output_response( + response=response, + guardrail_to_apply=guardrail, + ) + + assert guardrail.last_inputs is not None + assert guardrail.last_inputs["texts"] + assert "PONG" in result["result"]["parts"][0]["data"] + + +@pytest.mark.asyncio +async def test_process_input_messages_scans_data_parts(): + """The same bypass existed on the request/input side of the handler.""" + handler = A2AGuardrailHandler() + guardrail = MockGuardrail() + + data = { + "params": { + "message": { + "kind": "message", + "role": "user", + "parts": [{"kind": "data", "data": {"secret": "leak-me"}}], + } + } + } + + result = await handler.process_input_messages( + data=data, + guardrail_to_apply=guardrail, + ) + + assert guardrail.last_inputs is not None + assert any("leak-me" in t for t in guardrail.last_inputs["texts"]) + + data_part = result["params"]["message"]["parts"][0] + assert data_part["kind"] == "data" + assert "LEAK-ME" in data_part["data"]