diff --git a/litellm/llms/a2a/chat/guardrail_translation/handler.py b/litellm/llms/a2a/chat/guardrail_translation/handler.py index fbd1da749c2..c3c2e077ac8 100644 --- a/litellm/llms/a2a/chat/guardrail_translation/handler.py +++ b/litellm/llms/a2a/chat/guardrail_translation/handler.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index 84bbf6d20e1..334f38d982a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -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