test: address greptile comments

This commit is contained in:
Krrish Dholakia 2026-02-07 09:38:06 -08:00
parent 49580eb224
commit 6a524280d9
2 changed files with 88 additions and 71 deletions

View file

@ -217,12 +217,20 @@ class A2AGuardrailHandler(BaseTranslation):
"""
Process A2A streaming output by applying guardrails to accumulated text.
responses_so_far can be a list of JSON-RPC 2.0 objects (dict or NDJSON str), e.g.:
- task with history, status-update, artifact-update (with result.artifact.parts),
- then status-update (final). Text is extracted from result.artifact.parts,
result.message.parts, result.parts, etc., concatenated in order, guardrailed once,
then the combined guardrailed text is written into the first chunk that had text
and all other text parts in other chunks are cleared (in-place).
IN-PLACE MODIFICATION: This method mutates responses_so_far in place. Callers
must deep-copy the current chunk before calling if they intend to yield that
chunk, because we put the full guardrailed text in the first chunk and clear
all subsequent text parts to "".
Algorithm:
1. Parse each item (dict or NDJSON str) and collect text from result.artifact.parts,
result.message.parts, result.parts, etc.
2. Concatenate all texts in order, apply guardrail once to the combined string.
3. Write the full guardrailed text into the FIRST chunk's first text part.
4. Clear all other text parts in all chunks to "" (in-place).
responses_so_far: List of JSON-RPC 2.0 objects (dict or NDJSON str).
Returns: The same list (modified in place).
"""
from litellm.llms.a2a.common_utils import extract_text_from_a2a_response
@ -271,7 +279,12 @@ class A2AGuardrailHandler(BaseTranslation):
logging_obj=litellm_logging_obj,
)
guardrailed_texts = guardrailed_inputs.get("texts", [])
if not guardrailed_texts:
if not guardrailed_texts or len(guardrailed_texts) != 1:
# Guardrail should return exactly one text for combined input
verbose_proxy_logger.warning(
"A2A streaming guardrail returned unexpected texts count: %s",
len(guardrailed_texts or []),
)
return responses_so_far
guardrailed_text = guardrailed_texts[0]
@ -294,7 +307,7 @@ class A2AGuardrailHandler(BaseTranslation):
if not mappings:
continue
if orig_i == first_chunk_with_text:
# Put full guardrailed text in first text part; clear others
# Put full guardrailed text in first text part; clear others in this chunk
for task_idx, (path, part_idx) in enumerate(mappings):
text = guardrailed_text if task_idx == 0 else ""
self._apply_text_to_path(
@ -304,6 +317,7 @@ class A2AGuardrailHandler(BaseTranslation):
text=text,
)
else:
# Clear all text parts in non-first chunks (in-place)
for path, part_idx in mappings:
self._apply_text_to_path(
result=result,
@ -414,15 +428,21 @@ class A2AGuardrailHandler(BaseTranslation):
part_idx: int,
text: str,
) -> None:
"""Apply guardrailed text back to the specified path in the result."""
# Navigate to the parts list
current = result
"""Apply guardrailed text back to the specified path in the result (in-place)."""
current: Any = result
for key in path:
if key.isdigit():
# Array index
current = current[int(key)]
else:
current = current[key]
# Update the text in the part
current[part_idx]["text"] = text
if not isinstance(current, list) or part_idx >= len(current):
verbose_proxy_logger.warning(
"A2A _apply_text_to_path: invalid path or index path=%s part_idx=%s",
path,
part_idx,
)
return
part = current[part_idx]
if isinstance(part, dict) and "text" in part:
part["text"] = text

View file

@ -49,6 +49,52 @@ def _get_a2a_request_id(
return None
def _format_a2a_guardrail_error_chunk(
e: HTTPException,
call_type: Optional[str],
responses_so_far: List[Any],
request_data: dict,
) -> Optional[str]:
"""
Format HTTPException from guardrail as JSON-RPC 2.0 error chunk for A2A streaming.
When the response has already started, we cannot send an HTTP 4xx. For A2A (NDJSON)
streams, we yield an in-stream JSON-RPC error so the client receives the rejection.
Returns:
JSON-RPC error chunk string with trailing newline if call_type is A2A, else None.
Caller should yield the result when not None; otherwise re-raise the exception.
"""
if call_type is None or CallTypes(call_type) not in A2A_CALL_TYPES:
return None
request_id = _get_a2a_request_id(responses_so_far, request_data)
detail = (
e.detail
if isinstance(e.detail, dict)
else {"message": str(e.detail)}
)
return (
json.dumps(
{
"jsonrpc": "2.0",
"id": request_id,
"error": {
"code": -32603,
"message": detail.get(
"error", detail.get("message", str(e.detail))
),
"data": {
k: v
for k, v in detail.items()
if k not in ("error", "message")
},
},
}
)
+ "\n"
)
endpoint_guardrail_translation_mappings = None
@ -399,36 +445,10 @@ class UnifiedLLMGuardrails(CustomLogger):
user_api_key_dict=user_api_key_dict,
)
except HTTPException as e:
# Response already started (we already yielded chunks); cannot send 400.
# For A2A (NDJSON), yield an in-stream JSON-RPC error so the client sees it.
if call_type is not None and CallTypes(call_type) in A2A_CALL_TYPES:
request_id = _get_a2a_request_id(responses_so_far, request_data)
detail = (
e.detail
if isinstance(e.detail, dict)
else {"message": str(e.detail)}
)
error_chunk = (
json.dumps(
{
"jsonrpc": "2.0",
"id": request_id,
"error": {
"code": -32603,
"message": detail.get(
"error",
detail.get("message", str(e.detail)),
),
"data": {
k: v
for k, v in detail.items()
if k not in ("error", "message")
},
},
}
)
+ "\n"
)
error_chunk = _format_a2a_guardrail_error_chunk(
e, call_type, responses_so_far, request_data
)
if error_chunk is not None:
yield error_chunk
return
raise
@ -459,33 +479,10 @@ class UnifiedLLMGuardrails(CustomLogger):
user_api_key_dict=user_api_key_dict,
)
except HTTPException as e:
if call_type is not None and CallTypes(call_type) in A2A_CALL_TYPES:
request_id = _get_a2a_request_id(responses_so_far, request_data)
detail = (
e.detail
if isinstance(e.detail, dict)
else {"message": str(e.detail)}
)
error_chunk = (
json.dumps(
{
"jsonrpc": "2.0",
"id": request_id,
"error": {
"code": -32603,
"message": detail.get(
"error", detail.get("message", str(e.detail))
),
"data": {
k: v
for k, v in detail.items()
if k not in ("error", "message")
},
},
}
)
+ "\n"
)
error_chunk = _format_a2a_guardrail_error_chunk(
e, call_type, responses_so_far, request_data
)
if error_chunk is not None:
yield error_chunk
else:
raise