fix(guardrails): address Bedrock ApplyGuardrail chunking review feedback

Fixes three issues flagged in review of the chunking fallback: a single
oversized content item couldn't be split (only list-length bisection was
supported), a chunked request that got recovered still logged a stray
failure telemetry entry alongside the real outcome, and flattening chunk
outputs without positional bookkeeping could misalign masked text onto
the wrong original message once a chunk had nothing to mask.
This commit is contained in:
spencer-burridge 2026-07-28 13:18:57 -05:00
parent 7b7293eec3
commit 4271582892
2 changed files with 557 additions and 89 deletions

View file

@ -98,9 +98,6 @@ _BEDROCK_TOO_LARGE_ERROR_SUBSTRINGS = (
"too large",
"exceeds the maximum",
)
# Bisecting stops once a chunk is down to a single content item -- it cannot be
# split further, so a too-large error on it propagates as-is instead of looping.
_BEDROCK_APPLY_GUARDRAIL_MIN_CHUNK_SIZE = 1
# Exponential backoff for a chunk call throttled with ThrottlingException (429).
# Kept small: chunking already trades one oversized call for several smaller
# ones, so retries must not multiply per-request latency by an order of magnitude.
@ -153,6 +150,25 @@ class GuardrailMessageFilterResult(NamedTuple):
target_indices: Optional[List[int]]
class BedrockContentChunkResult(NamedTuple):
"""One chunk's ApplyGuardrail response, paired with enough bookkeeping to
reconstruct global masked-output positions once every chunk is back.
`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.
"""
response: BedrockGuardrailResponse
content: list[BedrockContentItem]
is_text_fragment: bool
class ApplyGuardrailMessageSelection(NamedTuple):
"""Messages selected for an apply_guardrail scan + write-back metadata."""
@ -810,37 +826,59 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
else:
event_type = GuardrailEventHooks.pre_call if source == "INPUT" else GuardrailEventHooks.post_call
content: List[BedrockContentItem] = bedrock_request_data.get("content") or []
content: list[BedrockContentItem] = bedrock_request_data.get("content") or []
# Contextual grounding scores the response holistically against the whole
# reference source; bisecting it would fragment that evaluation and produce
# misleading grounding scores, so a too-large error is never chunked here.
allow_chunking = not self._content_uses_contextual_grounding(content)
responses = await self._apply_guardrail_content_with_chunking(
content=content,
base_request_data=bedrock_request_data,
credentials=credentials,
aws_region_name=aws_region_name,
api_key=api_key,
try:
responses = await self._apply_guardrail_content_with_chunking(
content=content,
base_request_data=bedrock_request_data,
credentials=credentials,
aws_region_name=aws_region_name,
api_key=api_key,
request_data=request_data,
event_type=event_type,
start_time=start_time,
allow_chunking=allow_chunking,
)
except HTTPException as exc:
# A block is logged where it happens, inside _post_apply_guardrail_content,
# since chunking stops immediately and there is no later merged response to
# log instead. Anything else reaching here (unrecoverable too-large error,
# a non-size validation error, exhausted throttle retries) is a genuine
# end-to-end failure of this logical guardrail call and is logged once here.
if not isinstance(exc.detail, dict):
self._log_apply_guardrail_failure(
detail=exc.detail,
request_data=request_data,
event_type=event_type,
start_time=start_time,
)
raise
merged_response = self._merge_bedrock_guardrail_responses(responses)
self._log_apply_guardrail_success(
merged_response=merged_response,
request_data=request_data,
event_type=event_type,
start_time=start_time,
allow_chunking=allow_chunking,
)
return self._merge_bedrock_guardrail_responses(responses)
return merged_response
async def _apply_guardrail_content_with_chunking(
self,
content: List[BedrockContentItem],
content: list[BedrockContentItem],
base_request_data: dict,
credentials,
aws_region_name: str,
api_key: Optional[str],
api_key: str | None,
request_data: dict | None,
event_type: GuardrailEventHooks,
start_time: "datetime",
allow_chunking: bool,
) -> List[BedrockGuardrailResponse]:
) -> list[BedrockContentChunkResult]:
"""Post `content` to ApplyGuardrail, bisecting on a too-large error.
Tries `content` as a single call first. AWS's per-request "maximum input
@ -848,30 +886,41 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
predicted ahead of time, so it is only ever discovered reactively: on a 400
ValidationException whose message indicates the input was too large, the
content is split in half and each half is retried the same way (recursing
until every piece fits or cannot be split further). A real guardrail block
on any (sub-)chunk raises immediately -- callers must not lose that signal
by continuing to post the remaining chunks.
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
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
-- callers must not lose that signal by continuing to post the remaining
chunks.
"""
try:
response = await self._post_apply_guardrail_content_with_retry(
content=content,
base_request_data=base_request_data,
credentials=credentials,
aws_region_name=aws_region_name,
api_key=api_key,
request_data=request_data,
event_type=event_type,
start_time=start_time,
)
return [
await self._post_apply_guardrail_content_with_retry(
BedrockContentChunkResult(
response=response,
content=content,
base_request_data=base_request_data,
credentials=credentials,
aws_region_name=aws_region_name,
api_key=api_key,
request_data=request_data,
event_type=event_type,
start_time=start_time,
is_text_fragment=False,
)
]
except HTTPException as exc:
if (
allow_chunking
and len(content) > _BEDROCK_APPLY_GUARDRAIL_MIN_CHUNK_SIZE
and self._is_input_too_large_validation_error(exc.detail)
):
first_half, second_half = self._split_bedrock_content(content)
if allow_chunking and self._is_input_too_large_validation_error(exc.detail):
split_content = self._split_bedrock_content(content)
if split_content is None:
raise
first_half, second_half = split_content
is_text_fragment = len(content) == 1
first_results = await self._apply_guardrail_content_with_chunking(
content=first_half,
base_request_data=base_request_data,
@ -894,16 +943,19 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
start_time=start_time,
allow_chunking=allow_chunking,
)
return first_results + second_results
combined_results = first_results + second_results
if is_text_fragment:
return [result._replace(is_text_fragment=True) for result in combined_results]
return combined_results
raise
async def _post_apply_guardrail_content_with_retry(
self,
content: List[BedrockContentItem],
content: list[BedrockContentItem],
base_request_data: dict,
credentials,
aws_region_name: str,
api_key: Optional[str],
api_key: str | None,
request_data: dict | None,
event_type: GuardrailEventHooks,
start_time: "datetime",
@ -938,11 +990,11 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
async def _post_apply_guardrail_content(
self,
content: List[BedrockContentItem],
content: list[BedrockContentItem],
base_request_data: dict,
credentials,
aws_region_name: str,
api_key: Optional[str],
api_key: str | None,
request_data: dict | None,
event_type: GuardrailEventHooks,
start_time: "datetime",
@ -973,27 +1025,8 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
start_time=start_time,
)
#########################################################
# Add guardrail information to request trace
#########################################################
_json_response = httpx_response.json()
tracing_detail = self._build_tracing_detail(_json_response)
# Raw Bedrock JSON is passed here; match/regex redaction runs once inside
# CustomGuardrail.add_standard_logging_guardrail_information_to_request_data.
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_provider=self.guardrail_provider,
guardrail_json_response=_json_response,
request_data=request_data or {},
guardrail_status=self._get_bedrock_guardrail_response_status(response=httpx_response),
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,
tracing_detail=tracing_detail or None,
)
#########################################################
if httpx_response.status_code == 200:
_json_response = httpx_response.json()
# check if the response was flagged
verbose_proxy_logger.debug(
"Bedrock AI response : %s",
@ -1001,6 +1034,16 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
)
bedrock_guardrail_response = BedrockGuardrailResponse(**_json_response)
if self._should_raise_guardrail_blocked_exception(bedrock_guardrail_response):
# A block ends the whole chunking flow immediately (no further
# chunks are attempted), so it is logged here rather than by the
# caller -- there is no later "final merged response" to log instead.
self._log_apply_guardrail_attempt(
httpx_response=httpx_response,
json_response=_json_response,
request_data=request_data,
event_type=event_type,
start_time=start_time,
)
raise self._get_http_exception_for_blocked_guardrail(
bedrock_guardrail_response, request_data=request_data
)
@ -1014,8 +1057,78 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
)
raise HTTPException(status_code=status_code, detail=detail_message)
def _log_apply_guardrail_attempt(
self,
httpx_response: httpx.Response,
json_response: dict,
request_data: dict | None,
event_type: GuardrailEventHooks,
start_time: "datetime",
) -> None:
"""Log a single ApplyGuardrail HTTP attempt as-is (its own status,
derived from its own response). Used only for the blocked-content
case, which ends the whole chunking flow immediately."""
tracing_detail = self._build_tracing_detail(BedrockGuardrailResponse(**json_response))
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_provider=self.guardrail_provider,
guardrail_json_response=json_response,
request_data=request_data or {},
guardrail_status=self._get_bedrock_guardrail_response_status(response=httpx_response),
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,
tracing_detail=tracing_detail or None,
)
def _log_apply_guardrail_success(
self,
merged_response: BedrockGuardrailResponse,
request_data: dict | None,
event_type: GuardrailEventHooks,
start_time: "datetime",
) -> None:
"""Log one logical ApplyGuardrail call -- possibly several chunk calls
under the hood -- using its final merged response, so a chunked
request produces exactly one telemetry entry, the same as an
unchunked one would."""
tracing_detail = self._build_tracing_detail(merged_response)
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_provider=self.guardrail_provider,
guardrail_json_response=dict(merged_response),
request_data=request_data or {},
guardrail_status="success",
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,
tracing_detail=tracing_detail or None,
)
def _log_apply_guardrail_failure(
self,
detail: object,
request_data: dict | None,
event_type: GuardrailEventHooks,
start_time: "datetime",
) -> None:
"""Log one logical ApplyGuardrail call that failed end-to-end (an
unrecoverable too-large error, a non-size validation error, or
exhausted throttle retries) as a single failure, rather than logging
every failed attempt chunking made along the way."""
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_provider=self.guardrail_provider,
guardrail_json_response=str(detail),
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,
)
@staticmethod
def _content_uses_contextual_grounding(content: List[BedrockContentItem]) -> bool:
def _content_uses_contextual_grounding(content: list[BedrockContentItem]) -> bool:
"""True if any content item carries a contextual-grounding qualifier
(``grounding_source``, ``query``, or the ``guard_content`` the response
itself is tagged with once grounding is present)."""
@ -1026,15 +1139,36 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
@staticmethod
def _split_bedrock_content(
content: List[BedrockContentItem],
) -> Tuple[List[BedrockContentItem], List[BedrockContentItem]]:
content: list[BedrockContentItem],
) -> tuple[list[BedrockContentItem], list[BedrockContentItem]] | None:
"""Bisect `content` into two roughly-equal, non-empty halves.
Only called with ``len(content) > 1`` (guarded by the caller), so both
halves are always non-empty.
When `content` already holds more than one item, it is split by list
length. When it holds exactly one item, that item's own text is split
in half instead (a list of length 1 has no items left to bisect, but
one very long message is still a single content item). Returns None
when there is nothing left to split -- a single item whose text is
too short to halve into two non-empty pieces -- so the caller can
give up and propagate the original too-large error instead of
recursing forever.
"""
midpoint = max(1, len(content) // 2)
return content[:midpoint], content[midpoint:]
if len(content) > 1:
midpoint = max(1, len(content) // 2)
return content[:midpoint], content[midpoint:]
text_content = content[0].get("text") or BedrockTextContent()
text = text_content.get("text") or ""
if len(text) < 2:
return None
midpoint = len(text) // 2
qualifiers = text_content.get("qualifiers")
if qualifiers:
first_text = BedrockTextContent(text=text[:midpoint], qualifiers=qualifiers)
second_text = BedrockTextContent(text=text[midpoint:], qualifiers=qualifiers)
else:
first_text = BedrockTextContent(text=text[:midpoint])
second_text = BedrockTextContent(text=text[midpoint:])
return [BedrockContentItem(text=first_text)], [BedrockContentItem(text=second_text)]
@staticmethod
def _is_input_too_large_validation_error(detail: object) -> bool:
@ -1056,40 +1190,153 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
@staticmethod
def _merge_bedrock_guardrail_responses(
responses: List[BedrockGuardrailResponse],
chunk_results: list[BedrockContentChunkResult],
) -> BedrockGuardrailResponse:
"""Merge the per-chunk ApplyGuardrail responses of a chunked request into
one, so a caller cannot tell whether chunking happened.
Only ever called with responses that all passed (a block raises
immediately from ``_apply_guardrail_content_with_chunking`` and is never
added to this list), so ``action`` is included purely for completeness.
added to this list). ``action`` is only set on the merged response when
at least one chunk's raw response included it, and left absent otherwise
-- mirroring a real single-call response and matching what
``_build_tracing_detail`` treats as "Bedrock didn't report an action".
Per AWS's documented ApplyGuardrail contract, a single call's ``outputs``
is positionally parallel to the ``content`` items *of that call*: an
entry per item when anything in the call was masked, or an empty list
when nothing in the whole call was masked. Downstream masking
(``_apply_masking_to_messages``) walks the merged ``outputs`` by a single
running index across the *original, unchunked* message list, so a later
chunk's masked text must land at the same global position it would have
if chunking had never happened. Naively concatenating each chunk's
``outputs`` breaks that whenever a chunk had nothing masked (its empty
list would otherwise silently swallow its items' slots, shifting every
later chunk's masked text left onto the wrong message). So every
item -- masked or not -- always contributes exactly one entry here,
falling back to that item's own original (unmasked) text when its
chunk returned no output for it; a wholly-untouched result is then
collapsed back to an empty ``outputs`` list to match a real single-call
no-op response. A chunk that returns a nonzero output count not equal
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.
"""
if len(responses) == 1:
return responses[0]
logical_units = BedrockGuardrail._group_fragment_pairs(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)
merged_outputs: List[BedrockGuardrailOutput] = []
merged_assessments: List[dict] = []
merged_usage: Dict[str, Any] = {}
merged_action = "NONE"
for chunk_response in responses:
if chunk_response.get("action") == "GUARDRAIL_INTERVENED":
merged_action = "GUARDRAIL_INTERVENED"
merged_outputs.extend(chunk_response.get("outputs") or chunk_response.get("output") or [])
merged_assessments.extend(chunk_response.get("assessments") or [])
for key, value in (chunk_response.get("usage") or {}).items():
if isinstance(value, (int, float)):
merged_usage[key] = merged_usage.get(key, 0) + value
actions = tuple(
chunk_result.response.get("action")
for chunk_result in chunk_results
if isinstance(chunk_result.response.get("action"), str)
)
merged_action = (
"GUARDRAIL_INTERVENED" if "GUARDRAIL_INTERVENED" in actions else (actions[-1] if actions else None)
)
merged_assessments = [
assessment
for chunk_result in chunk_results
for assessment in (chunk_result.response.get("assessments") or [])
]
any_usage_reported = any(chunk_result.response.get("usage") for chunk_result in chunk_results)
merged: BedrockGuardrailResponse = BedrockGuardrailResponse(action=merged_action)
if merged_outputs:
merged: BedrockGuardrailResponse = BedrockGuardrailResponse()
if merged_action is not None:
merged["action"] = merged_action
if merged_outputs and any_masked:
merged["outputs"] = merged_outputs
if merged_assessments:
merged["assessments"] = merged_assessments
if merged_usage:
merged["usage"] = cast(BedrockGuardrailUsage, merged_usage)
if any_usage_reported:
merged["usage"] = BedrockGuardrail._sum_bedrock_guardrail_usage(chunk_results)
return merged
@staticmethod
def _sum_bedrock_guardrail_usage(
chunk_results: list[BedrockContentChunkResult],
) -> BedrockGuardrailUsage:
"""Sum each chunk's ``usage`` counters field-by-field into one totals dict."""
chunk_usages = tuple(chunk_result.response.get("usage") or {} for chunk_result in chunk_results)
def total(key: str) -> int:
return sum(usage.get(key) or 0 for usage in chunk_usages)
return BedrockGuardrailUsage(
topicPolicyUnits=total("topicPolicyUnits"),
contentPolicyUnits=total("contentPolicyUnits"),
wordPolicyUnits=total("wordPolicyUnits"),
sensitiveInformationPolicyUnits=total("sensitiveInformationPolicyUnits"),
sensitiveInformationPolicyFreeUnits=total("sensitiveInformationPolicyFreeUnits"),
contextualGroundingPolicyUnits=total("contextualGroundingPolicyUnits"),
)
@staticmethod
def _group_fragment_pairs(
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``."""
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
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.
Per AWS's documented ApplyGuardrail contract, a single call's
``outputs`` is positionally parallel to the ``content`` items *of that
call*: an entry per item when anything in the call was masked, or an
empty list when nothing in the whole call was masked. Downstream
masking (``_apply_masking_to_messages``) walks the merged ``outputs``
by a single running index across the *original, unchunked* message
list, so a later chunk's masked text must land at the same global
position it would have if chunking had never happened. So every item
-- masked or not -- always contributes exactly one entry here, falling
back to that item's own original (unmasked) text when its chunk
returned no output for it. A chunk that returns a nonzero output count
not equal 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.
"""
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)
(chunk_result,) = unit
chunk_outputs = chunk_result.response.get("outputs") or chunk_result.response.get("output") or []
if len(chunk_outputs) == len(chunk_result.content):
return list(chunk_outputs), bool(chunk_outputs)
if not chunk_outputs:
return [
BedrockGuardrailOutput(text=(item.get("text") or {}).get("text") or "") for item in chunk_result.content
], False
return list(chunk_outputs), True
async def _sign_and_post(
self,
prepared_request: "AWSPreparedRequest",

View file

@ -3537,13 +3537,13 @@ async def test_apply_guardrail_does_not_chunk_on_non_size_validation_error():
@pytest.mark.asyncio
async def test_apply_guardrail_too_large_on_single_item_propagates_original_error():
"""A too-large error on content that is already down to a single item
cannot be bisected further; the original error must propagate rather than
looping or crashing."""
async def test_apply_guardrail_too_large_on_unsplittable_text_propagates_original_error():
"""A too-large error on content that has been bisected down to text too
short to split further (< 2 characters) must propagate the original error
rather than looping or crashing."""
guardrail = _bedrock_guardrail_for_chunk_tests()
messages = [{"role": "user", "content": "one giant single block of text"}]
messages = [{"role": "user", "content": "a"}]
mock_credentials = MagicMock()
mock_credentials.access_key = "k"
@ -3568,6 +3568,46 @@ async def test_apply_guardrail_too_large_on_single_item_propagates_original_erro
assert exc_info.value.status_code == 400
@pytest.mark.asyncio
async def test_apply_guardrail_too_large_on_single_item_splits_by_text_and_succeeds():
"""A too-large error on content that is already down to a single content
item must be bisected by that item's own text (not abandoned), so an
oversized single message can still be scanned successfully in halves."""
guardrail = _bedrock_guardrail_for_chunk_tests()
messages = [{"role": "user", "content": "one giant single block of 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:
return _too_large_validation_httpx_response()
return _passing_bedrock_httpx_response(f"half-{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()),
):
mock_post.side_effect = _post_side_effect
response = await guardrail.make_bedrock_api_request(
source="INPUT",
messages=messages,
request_data={"model": "bedrock-nova-micro"},
)
assert mock_post.await_count == 3
assert response.get("action") == "NONE"
@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
@ -3619,3 +3659,184 @@ async def test_apply_guardrail_chunk_retries_after_throttling_then_succeeds():
mock_sleep.assert_awaited()
output_texts = [o.get("text") for o in result.get("outputs") or []]
assert output_texts == ["chunk-1", "chunk-2"]
@pytest.mark.asyncio
async def test_apply_guardrail_chunking_logs_exactly_once_as_success():
"""A too-large 400 that is recovered by chunking must not leave behind a
'guardrail_failed_to_respond' telemetry entry for the initial oversized
attempt: the whole logical request (1 too-large attempt + 2 chunk
attempts here) must produce exactly one standard-logging entry, and it
must reflect the eventual success, not the transient too-large failure."""
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:
return _too_large_validation_httpx_response()
if call_count == 2:
return _passing_bedrock_httpx_response("chunk-1")
return _passing_bedrock_httpx_response("chunk-2")
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
await guardrail.make_bedrock_api_request(
source="INPUT",
messages=messages,
request_data={"model": "bedrock-nova-micro"},
)
assert call_count == 3
mock_log.assert_called_once()
assert mock_log.call_args.kwargs["guardrail_status"] == "success"
@pytest.mark.asyncio
async def test_apply_guardrail_unrecoverable_failure_logs_exactly_once_as_failed():
"""A too-large error that cannot be recovered (chunking disabled by
contextual grounding) must still log exactly once, as a failure -- not be
silently dropped by the chunking telemetry consolidation."""
guardrail = _bedrock_guardrail_for_chunk_tests()
messages = [
{
"role": "system",
"content": [{"type": "grounding_source", "text": "reference source text"}],
},
{"role": "user", "content": "what does the source say?"},
]
model_response = ModelResponse()
model_response.choices = [
litellm.Choices(message=litellm.Message(content="a grounded answer", role="assistant"))
]
mock_credentials = MagicMock()
mock_credentials.access_key = "k"
mock_credentials.secret_key = "s"
mock_credentials.token = None
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.return_value = _too_large_validation_httpx_response()
with pytest.raises(HTTPException):
await guardrail.make_bedrock_api_request(
source="OUTPUT",
messages=messages,
response=model_response,
request_data={"model": "bedrock-nova-micro"},
)
mock_post.assert_awaited_once()
mock_log.assert_called_once()
assert mock_log.call_args.kwargs["guardrail_status"] == "guardrail_failed_to_respond"
@pytest.mark.asyncio
async def test_apply_guardrail_chunk_merge_preserves_masking_position():
"""An earlier chunk that comes back clean (empty `outputs`) must not
shift a later chunk's masked text onto the wrong message. Regression for:
flattening outputs without positional metadata let a later chunk's PII
redaction get applied to the first message while the actual PII-bearing
message (in a later chunk) was forwarded unmasked."""
guardrail = BedrockGuardrail(
guardrail_name="test-bedrock-guard",
guardrailIdentifier="test-guardrail",
guardrailVersion="DRAFT",
disable_exception_on_block=False,
mask_request_content=True,
)
request_data = {
"model": "bedrock-nova-micro",
"messages": [
{"role": "user", "content": "clean chunk with nothing to mask"},
{"role": "user", "content": "chunk with PII: John Doe"},
],
}
mock_credentials = MagicMock()
mock_credentials.access_key = "k"
mock_credentials.secret_key = "s"
mock_credentials.token = None
call_count = 0
def _clean_httpx_response() -> MagicMock:
response = MagicMock()
response.status_code = 200
response.json.return_value = {"action": "NONE", "assessments": []}
return response
def _masked_httpx_response() -> MagicMock:
response = MagicMock()
response.status_code = 200
response.json.return_value = {
"action": "GUARDRAIL_INTERVENED",
"outputs": [{"text": "chunk with PII: [NAME]"}],
"assessments": [
{
"sensitiveInformationPolicy": {
"piiEntities": [{"type": "NAME", "match": "John Doe", "action": "ANONYMIZED"}]
}
}
],
}
return response
async def _post_side_effect(*_args, **_kwargs):
nonlocal call_count
call_count += 1
if call_count == 1:
return _too_large_validation_httpx_response()
if call_count == 2:
return _clean_httpx_response()
return _masked_httpx_response()
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
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(),
cache=DualCache(),
data=request_data,
call_type="acompletion",
)
assert call_count == 3
updated_messages = request_data["messages"]
assert updated_messages[0]["content"] == "clean chunk with nothing to mask"
assert updated_messages[1]["content"] == "chunk with PII: [NAME]"