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).
This commit is contained in:
STiFLeR7 2026-07-30 11:00:35 +05:30
parent 8cb053ff3f
commit 721babb0e7
No known key found for this signature in database
5 changed files with 191 additions and 26 deletions

View file

@ -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

View file

@ -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)

View file

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