diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 8bfe5027b77..5c9baf21fff 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -1243,90 +1243,104 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ] = stream_chunk_builder( chunks=all_chunks, ) - if isinstance(assembled_model_response, ModelResponse): - #################################################################### - ########## 1. Make Bedrock Apply Guardrail API requests ########## + try: + if isinstance(assembled_model_response, ModelResponse): + #################################################################### + ########## 1. Make Bedrock Apply Guardrail API requests ########## - # Bedrock will raise an exception if this violates the guardrail policy - ################################################################### - # Determine if INPUT validation is needed in post_call - # Skip INPUT validation if pre_call or during_call is already enabled - # (to avoid redundant validation - those hooks would have already validated INPUT) - should_validate_input = not ( - self._event_hook_is_event_type(GuardrailEventHooks.pre_call) - or self._event_hook_is_event_type(GuardrailEventHooks.during_call) - ) - - output_guardrail_response: Optional[ - Union[BedrockGuardrailResponse, str] - ] = None - - if should_validate_input: - # Create tasks for parallel execution - input_filter = self._prepare_guardrail_messages_for_role( - messages=request_data.get("messages") + # Bedrock will raise an exception if this violates the guardrail policy + ################################################################### + # Determine if INPUT validation is needed in post_call + # Skip INPUT validation if pre_call or during_call is already enabled + # (to avoid redundant validation - those hooks would have already validated INPUT) + should_validate_input = not ( + self._event_hook_is_event_type(GuardrailEventHooks.pre_call) + or self._event_hook_is_event_type(GuardrailEventHooks.during_call) ) - input_messages = input_filter.payload_messages or request_data.get( - "messages" - ) - input_task = self.make_bedrock_api_request( - source="INPUT", - messages=input_messages, - request_data=request_data, - logging_event_type=GuardrailEventHooks.post_call, - ) # Only input messages - output_task = self.make_bedrock_api_request( - source="OUTPUT", - response=assembled_model_response, - request_data=request_data, - logging_event_type=GuardrailEventHooks.post_call, - ) # Only response - # Execute both requests in parallel - try: - _, output_guardrail_response = await asyncio.gather( - input_task, output_task + output_guardrail_response: Optional[ + Union[BedrockGuardrailResponse, str] + ] = None + + if should_validate_input: + # Create tasks for parallel execution + input_filter = self._prepare_guardrail_messages_for_role( + messages=request_data.get("messages") ) - except GuardrailInterventionNormalStringError as e: - output_guardrail_response = e.message - else: - # Only run OUTPUT validation (INPUT was already validated in pre_call or during_call) - try: - output_guardrail_response = await self.make_bedrock_api_request( + input_messages = input_filter.payload_messages or request_data.get( + "messages" + ) + input_task = self.make_bedrock_api_request( + source="INPUT", + messages=input_messages, + request_data=request_data, + logging_event_type=GuardrailEventHooks.post_call, + ) # Only input messages + output_task = self.make_bedrock_api_request( source="OUTPUT", response=assembled_model_response, request_data=request_data, logging_event_type=GuardrailEventHooks.post_call, + ) # Only response + + # Execute both requests in parallel + try: + _, output_guardrail_response = await asyncio.gather( + input_task, output_task + ) + except GuardrailInterventionNormalStringError as e: + output_guardrail_response = e.message + else: + # Only run OUTPUT validation (INPUT was already validated in pre_call or during_call) + try: + output_guardrail_response = await self.make_bedrock_api_request( + source="OUTPUT", + response=assembled_model_response, + request_data=request_data, + logging_event_type=GuardrailEventHooks.post_call, + ) + except GuardrailInterventionNormalStringError as e: + output_guardrail_response = e.message + + ######################################################################### + ########## 2. Apply masking to response with output guardrail response ########## + ######################################################################### + if isinstance(output_guardrail_response, str): + assembled_model_response = self.create_guardrail_blocked_response( + response=output_guardrail_response + ) + elif output_guardrail_response is not None: + self._apply_masking_to_response( + response=assembled_model_response, + bedrock_guardrail_response=output_guardrail_response, ) - except GuardrailInterventionNormalStringError as e: - output_guardrail_response = e.message - ######################################################################### - ########## 2. Apply masking to response with output guardrail response ########## - ######################################################################### - if isinstance(output_guardrail_response, str): - assembled_model_response = self.create_guardrail_blocked_response( - response=output_guardrail_response - ) - elif output_guardrail_response is not None: - self._apply_masking_to_response( - response=assembled_model_response, - bedrock_guardrail_response=output_guardrail_response, + ######################################################################### + ########## 3. Return the (potentially masked) chunks ########## + ######################################################################### + mock_response = MockResponseIterator( + model_response=assembled_model_response ) - ######################################################################### - ########## 3. Return the (potentially masked) chunks ########## - ######################################################################### - mock_response = MockResponseIterator( - model_response=assembled_model_response - ) - - # Return the reconstructed stream - async for chunk in mock_response: - yield chunk - else: - for chunk in all_chunks: - yield chunk + # Return the reconstructed stream + async for chunk in mock_response: + yield chunk + else: + for chunk in all_chunks: + yield chunk + except HTTPException as e: + # An async generator cannot propagate an exception as an HTTP error response + # after the 200 header has been written. Convert hard-block exceptions to + # in-band SSE error frames so the client receives a structured error. + error_data = { + "error": { + "message": str(e.detail), + "code": e.status_code, + "type": "guardrail_violation", + } + } + yield f"data: {json.dumps(error_data)}\n\n" # type: ignore + yield "data: [DONE]\n\n" # type: ignore def _extract_masked_texts_from_response( self, bedrock_guardrail_response: BedrockGuardrailResponse