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"]