mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
8cb053ff3f
commit
721babb0e7
5 changed files with 191 additions and 26 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
0
tests/test_litellm/llms/a2a/chat/__init__.py
Normal file
0
tests/test_litellm/llms/a2a/chat/__init__.py
Normal 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"]
|
||||
Loading…
Add table
Reference in a new issue