fix(guardrails): group all fragments of one item and stop double-logging

Two defects found in review, both invisible to the existing tests.

Fragment grouping assumed a split content item always produces exactly two
adjacent fragments. That holds for one bisection level but not two: an item
split twice yields four fragments, which were regrouped in fixed pairs into
two output entries for a single message. Since masking walks the merged
outputs by a running index across the original, unchunked message list, that
message was written back truncated to its first half and every later message
shifted. Fragments now carry the size of the group they belong to, so any
number of them collapse back into exactly one output entry.

Telemetry was also double-counted. AsyncHTTPHandler.post calls
raise_for_status(), so every non-200 from Bedrock reaches _sign_and_post's
error path, which logged guardrail_failed_to_respond before re-raising as an
HTTPException that the consolidating caller then logged again. A request
recovered by chunking reported one failure per rejected attempt plus a
success. The ApplyGuardrail path now opts out of that per-attempt logging,
since it owns consolidated per-request logging; the connection-level branch
still logs, as nothing else records it.

The existing tests missed both because their mocks return a non-200 response
object, while the real client raises. Added a helper that raises a genuine
httpx.HTTPStatusError so these paths are covered the way production hits
them, plus a case asserting an unrecoverable failure still logs exactly once
rather than zero times.
This commit is contained in:
spencer-burridge 2026-07-29 13:08:21 -05:00
parent 71d99c7931
commit 2b449f0cd0
2 changed files with 257 additions and 50 deletions

View file

@ -179,16 +179,20 @@ class BedrockContentChunkResult(NamedTuple):
`content` is the exact content items this chunk was called with -- needed
so an all-clear chunk (empty `outputs`) can still contribute one unmasked
placeholder per item it covers, keeping every later chunk's masked text
aligned to its original global position. `is_text_fragment` is True when
this chunk is one half of a single content item's own text (split because
a list of length 1 could not be bisected by list length) -- its sibling
fragment must be concatenated back into that one item's masked output,
not treated as a second item.
aligned to its original global position. `fragment_group_size` is 1 for an
ordinary chunk, and otherwise the total number of consecutive chunk results
that together make up ONE original content item's own text (split because a
list of length 1 could not be bisected by list length). All of them must be
concatenated back into that one item's masked output rather than treated as
separate items. It is a count rather than a boolean because one item can be
bisected more than once: two levels of splitting produce four fragments for
a single item, not two, and grouping them in fixed pairs would emit two
outputs for one message and shift every later message's masked text.
"""
response: BedrockGuardrailResponse
content: list[BedrockContentItem]
is_text_fragment: bool
fragment_group_size: int
class ApplyGuardrailMessageSelection(NamedTuple):
@ -922,10 +926,12 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
until every piece fits or cannot be split further). A single oversized
content item (one very long message) is split by its own text instead of
by list length, since a list of length 1 has no items left to bisect --
the two text fragments are tagged ``is_text_fragment=True`` so the merge
the resulting fragments all carry a ``fragment_group_size`` so the merge
step can recombine them into the one content item they came from, rather
than treating each fragment as its own item when reconstructing positions
for masking. A real guardrail block on any (sub-)chunk raises immediately
for masking. That count covers however many fragments the item ended up
split into, not just two, since it can be bisected repeatedly. A real
guardrail block on any (sub-)chunk raises immediately
-- callers must not lose that signal by continuing to post the remaining
chunks.
"""
@ -944,7 +950,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
BedrockContentChunkResult(
response=response,
content=content,
is_text_fragment=False,
fragment_group_size=1,
)
]
except HTTPException as exc:
@ -953,7 +959,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
if split_content is None:
raise
first_half, second_half = split_content
is_text_fragment = len(content) == 1
is_single_item_text_split = len(content) == 1
verbose_proxy_logger.warning(
"Bedrock Guardrail: ApplyGuardrail rejected %d content item(s) as too large; "
"splitting into %d + %d and retrying each",
@ -984,8 +990,13 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
allow_chunking=allow_chunking,
)
combined_results = first_results + second_results
if is_text_fragment:
return [result._replace(is_text_fragment=True) for result in combined_results]
if is_single_item_text_split:
# Every leaf below this point came from one content item's own
# text, however many levels deep the splitting went. Stamping
# the total count on all of them (overwriting any smaller count
# an inner split set) is what lets the merge step regroup them
# into exactly one output entry for that one item.
return [result._replace(fragment_group_size=len(combined_results)) for result in combined_results]
return combined_results
raise
@ -1074,6 +1085,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
request_data=request_data,
event_type=event_type,
start_time=start_time,
log_transport_failure=False,
)
if httpx_response.status_code == 200:
@ -1352,7 +1364,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
to its item count is passed through as-is instead of guessed at, since
AWS's docs don't cover partial masking within one multi-item call.
"""
logical_units = BedrockGuardrail._group_fragment_pairs(chunk_results)
logical_units = BedrockGuardrail._group_fragment_units(chunk_results)
per_unit_outputs = tuple(BedrockGuardrail._merge_logical_unit_outputs(unit) for unit in logical_units)
merged_outputs = [output for outputs, _ in per_unit_outputs for output in outputs]
any_masked = any(masked for _, masked in per_unit_outputs)
@ -1403,32 +1415,33 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
)
@staticmethod
def _group_fragment_pairs(
def _group_fragment_units(
chunk_results: list[BedrockContentChunkResult],
) -> list[tuple[BedrockContentChunkResult, ...]]:
"""Group consecutive text-fragment chunk results into the sibling pairs
that originated from one content item's own text, leaving every other
chunk result as a unit of one. Fragments are always produced (and thus
appear here) as adjacent sibling pairs -- see ``_split_bedrock_content``."""
"""Group consecutive text-fragment chunk results back into the one content
item each group came from, leaving every ordinary chunk result as a unit of
one.
The group size is read off the results themselves rather than assumed,
because a single content item can be bisected repeatedly: two levels of
splitting yield four fragments for one item, not two. Assuming a fixed pair
here would emit two outputs for one message and shift every later message's
masked text onto the wrong message."""
units: list[tuple[BedrockContentChunkResult, ...]] = []
index = 0
while index < len(chunk_results):
chunk_result = chunk_results[index]
if chunk_result.is_text_fragment:
units.append((chunk_result, chunk_results[index + 1]))
index += 2
else:
units.append((chunk_result,))
index += 1
span = max(1, chunk_results[index].fragment_group_size)
units.append(tuple(chunk_results[index : index + span]))
index += span
return units
@staticmethod
def _merge_logical_unit_outputs(
unit: tuple[BedrockContentChunkResult, ...],
) -> tuple[list[BedrockGuardrailOutput], bool]:
"""Reduce one logical unit (a fragment pair or a single chunk result)
to the ``BedrockGuardrailOutput`` entries it contributes to the merged
response, plus whether any masking actually happened in it.
"""Reduce one logical unit (a fragment group of any size, or a single chunk
result) to the ``BedrockGuardrailOutput`` entries it contributes to the
merged response, plus whether any masking actually happened in it.
Per AWS's documented ApplyGuardrail contract, a single call's
``outputs`` is positionally parallel to the ``content`` items *of that
@ -1445,18 +1458,23 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
at, since AWS's docs don't cover partial masking within one multi-item
call.
"""
if len(unit) == 2:
first, second = unit
first_outputs = first.response.get("outputs") or first.response.get("output") or []
second_outputs = second.response.get("outputs") or second.response.get("output") or []
first_source = (first.content[0].get("text") or {}).get("text") or ""
second_source = (second.content[0].get("text") or {}).get("text") or ""
first_text = first_outputs[0].get("text") if first_outputs else first_source
second_text = second_outputs[0].get("text") if second_outputs else second_source
merged_text = (first_text if first_text is not None else first_source) + (
second_text if second_text is not None else second_source
)
return [BedrockGuardrailOutput(text=merged_text)], bool(first_outputs or second_outputs)
if len(unit) > 1:
# Every result here is one fragment of a single content item's text, so
# the whole group collapses to one entry: each fragment's masked text
# (or its own original text when that fragment came back unmasked),
# concatenated in order. Holds for any group size, not just two.
def fragment_outputs(result: BedrockContentChunkResult) -> list[BedrockGuardrailOutput]:
return list(result.response.get("outputs") or result.response.get("output") or [])
def fragment_text(result: BedrockContentChunkResult) -> str:
source = (result.content[0].get("text") or {}).get("text") or ""
outputs = fragment_outputs(result)
masked = outputs[0].get("text") if outputs else None
return masked if masked is not None else source
merged_text = "".join(fragment_text(result) for result in unit)
any_masked = any(fragment_outputs(result) for result in unit)
return [BedrockGuardrailOutput(text=merged_text)], any_masked
(chunk_result,) = unit
chunk_outputs = chunk_result.response.get("outputs") or chunk_result.response.get("output") or []
@ -1474,6 +1492,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
request_data: dict | None,
event_type: GuardrailEventHooks,
start_time: "datetime",
log_transport_failure: bool = True,
) -> httpx.Response:
"""POST a signed Bedrock request, logging+raising on network/HTTP errors.
@ -1481,6 +1500,20 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
transport-error handling cannot drift. Returns the raw ``httpx.Response`` on
success (including non-2xx that httpx did not raise on); the 200-path logging,
status and tracing stay with each caller because the two APIs report differently.
``log_transport_failure=False`` suppresses the ``guardrail_failed_to_respond``
entry for a non-200 that is re-raised as an ``HTTPException``, for callers that
own consolidated per-request logging. The ApplyGuardrail path needs this:
``AsyncHTTPHandler.post`` calls ``raise_for_status()``, so every non-200 lands
in this handler, and one logical request can legitimately produce several of
them (a too-large probe, then each rejected bisection level) while still
succeeding overall. Logging per attempt would report a recovered request as
several failures plus a success.
The connection-level branch below (timeout, endpoint down) still logs
unconditionally: it re-raises the original exception rather than an
``HTTPException``, so no consolidating caller catches it, and suppressing it
would drop the only record of the failure.
"""
try:
return await self.async_handler.post(
@ -1501,16 +1534,17 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
status_code,
detail_message,
) = self._parse_bedrock_guardrail_error_response(err_response)
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_provider=self.guardrail_provider,
guardrail_json_response={"error": detail_message},
request_data=request_data or {},
guardrail_status="guardrail_failed_to_respond",
start_time=start_time.timestamp(),
end_time=datetime.now(timezone.utc).timestamp(),
duration=(datetime.now(timezone.utc) - start_time).total_seconds(),
event_type=event_type,
)
if log_transport_failure:
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_provider=self.guardrail_provider,
guardrail_json_response={"error": detail_message},
request_data=request_data or {},
guardrail_status="guardrail_failed_to_respond",
start_time=start_time.timestamp(),
end_time=datetime.now(timezone.utc).timestamp(),
duration=(datetime.now(timezone.utc) - start_time).total_seconds(),
event_type=event_type,
)
raise HTTPException(status_code=status_code, detail=detail_message) from e
except HTTPException:
raise

View file

@ -7,6 +7,7 @@ import os
import sys
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from fastapi import HTTPException
@ -3688,6 +3689,178 @@ async def test_apply_guardrail_too_large_on_single_item_splits_by_text_and_succe
assert response.get("action") == "NONE"
def _raised_bedrock_error(status_code: int, message: str) -> httpx.HTTPStatusError:
"""A non-200 the way `AsyncHTTPHandler.post` actually surfaces it.
That handler calls `response.raise_for_status()`, so in production a non-200 from
Bedrock arrives as a raised `httpx.HTTPStatusError` carrying the response, never
as a returned response object. Tests that return the response instead exercise a
branch real traffic never reaches. A real `httpx.Response` is used rather than a
MagicMock because the transport helper branches on
`isinstance(err_response, httpx.Response)`."""
response = httpx.Response(
status_code=status_code,
json={"message": message},
request=httpx.Request("POST", "https://bedrock-runtime.us-east-1.amazonaws.com/guardrail"),
)
return httpx.HTTPStatusError(message, request=response.request, response=response)
_TOO_LARGE_MESSAGE = "Input is too long. Content size exceeds the maximum input size in text units."
@pytest.mark.asyncio
async def test_apply_guardrail_chunking_logs_once_when_client_raises_for_status():
"""The too-large attempt recovered by chunking must still produce exactly one
telemetry entry when the HTTP client raises for status, which is what really
happens: `AsyncHTTPHandler.post` calls `raise_for_status()`.
Regression for per-attempt `guardrail_failed_to_respond` entries leaking out of
the transport helper on a request that ultimately succeeded, which made a
recovered request look like several failures plus a success."""
guardrail = _bedrock_guardrail_for_chunk_tests()
messages = [
{"role": "user", "content": "chunk one text"},
{"role": "user", "content": "chunk two text"},
]
mock_credentials = MagicMock()
mock_credentials.access_key = "k"
mock_credentials.secret_key = "s"
mock_credentials.token = None
call_count = 0
async def _post_side_effect(*_args, **_kwargs):
nonlocal call_count
call_count += 1
if call_count == 1:
raise _raised_bedrock_error(400, _TOO_LARGE_MESSAGE)
return _passing_bedrock_httpx_response(f"chunk-{call_count}")
with (
patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post,
patch.object(guardrail, "_load_credentials", return_value=(mock_credentials, "us-east-1")),
patch.object(guardrail, "_prepare_request", return_value=MagicMock()),
patch.object(
guardrail,
"add_standard_logging_guardrail_information_to_request_data",
) as mock_log,
):
mock_post.side_effect = _post_side_effect
result = await guardrail.make_bedrock_api_request(
source="INPUT",
messages=messages,
request_data={"model": "bedrock-nova-micro"},
)
assert call_count == 3
assert result.get("action") == "NONE"
statuses = [call.kwargs.get("guardrail_status") for call in mock_log.call_args_list]
assert statuses == ["success"]
@pytest.mark.asyncio
async def test_apply_guardrail_unrecoverable_failure_still_logs_once_when_client_raises():
"""Suppressing the transport helper's per-attempt logging must not swallow the only
record of a genuine failure: an unsplittable too-large request still has to produce
exactly one `guardrail_failed_to_respond` entry, not zero."""
guardrail = _bedrock_guardrail_for_chunk_tests()
# Single character: `_split_bedrock_content` cannot halve this into two non-empty
# pieces, so chunking gives up and the original error propagates.
messages = [{"role": "user", "content": "x"}]
mock_credentials = MagicMock()
mock_credentials.access_key = "k"
mock_credentials.secret_key = "s"
mock_credentials.token = None
async def _post_side_effect(*_args, **_kwargs):
raise _raised_bedrock_error(400, _TOO_LARGE_MESSAGE)
with (
patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post,
patch.object(guardrail, "_load_credentials", return_value=(mock_credentials, "us-east-1")),
patch.object(guardrail, "_prepare_request", return_value=MagicMock()),
patch.object(
guardrail,
"add_standard_logging_guardrail_information_to_request_data",
) as mock_log,
):
mock_post.side_effect = _post_side_effect
with pytest.raises(HTTPException):
await guardrail.make_bedrock_api_request(
source="INPUT",
messages=messages,
request_data={"model": "bedrock-nova-micro"},
)
statuses = [call.kwargs.get("guardrail_status") for call in mock_log.call_args_list]
assert statuses == ["guardrail_failed_to_respond"]
@pytest.mark.asyncio
async def test_apply_guardrail_single_item_split_twice_still_yields_one_output_per_item():
"""One oversized content item that needs two levels of text bisection ends up
as four text fragments, and all four must still collapse back into exactly
ONE output entry, because they all came from one original content item.
Downstream masking (`_apply_masking_to_messages`) walks the merged outputs by
a running index across the original, unchunked message list, so emitting more
than one entry for a single message shifts every later message's masked text
onto the wrong message and drops the surplus. Regression for fragment
grouping assuming fragments only ever arrive as adjacent sibling *pairs*,
which holds for one bisection level but not for two."""
guardrail = _bedrock_guardrail_for_chunk_tests()
messages = [{"role": "user", "content": "aaaa bbbb cccc dddd eeee ffff gggg hhhh"}]
mock_credentials = MagicMock()
mock_credentials.access_key = "k"
mock_credentials.secret_key = "s"
mock_credentials.token = None
# whole item -> [first half] -> [q1] [q2] -> [second half] -> [q3] [q4].
# Only the four quarters fit; the whole item and both halves are too large.
responses = [
_too_large_validation_httpx_response(), # whole single item
_too_large_validation_httpx_response(), # first half
_passing_bedrock_httpx_response("q1"),
_passing_bedrock_httpx_response("q2"),
_too_large_validation_httpx_response(), # second half
_passing_bedrock_httpx_response("q3"),
_passing_bedrock_httpx_response("q4"),
]
async def _post_side_effect(*_args, **_kwargs):
return responses.pop(0)
with (
patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post,
patch.object(guardrail, "_load_credentials", return_value=(mock_credentials, "us-east-1")),
patch.object(guardrail, "_prepare_request", return_value=MagicMock()),
):
mock_post.side_effect = _post_side_effect
result = await guardrail.make_bedrock_api_request(
source="INPUT",
messages=messages,
request_data={"model": "bedrock-nova-micro"},
)
assert mock_post.await_count == 7
assert not responses
assert result.get("action") == "NONE"
output_texts = [o.get("text") for o in result.get("outputs") or []]
# One original content item in, so exactly one output entry out, carrying all
# four fragments' text in order.
assert output_texts == ["q1q2q3q4"]
@pytest.mark.asyncio
async def test_apply_guardrail_chunk_retries_after_throttling_then_succeeds():
"""A chunk call throttled with a 429 must be retried with backoff and