fix: convert HTTPException to SSE error frames in bedrock guardrail streaming hook

This commit is contained in:
Nicholas Couture 2026-04-24 15:01:59 +10:00
parent 0992bf2271
commit c90b0b2617
No known key found for this signature in database

View file

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