diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index c924aa2ca54..8e52aca1f43 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -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", diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py index c9308de2e5b..f9ffb44d6fa 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py @@ -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]"