mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
7b7293eec3
commit
4271582892
2 changed files with 557 additions and 89 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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]"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue