fix(guardrails): hybrid bin-pack+bisection chunking, whitespace-safe splits

Rework Bedrock ApplyGuardrail chunking from pure reactive bisection to a
hybrid strategy: bin-pack content into fixed-budget batches up front as
the fast path, falling back to the existing recursive bisection only for
a batch AWS still rejects as too large. Avoids paying O(log n) round
trips on every oversized request when a single pass would do.

Also switch single-item text splitting from a raw character midpoint to
the nearest whitespace boundary, so a fragment never starts or ends
mid-word. Closes the accidental-severing case from review; the residual
gap (a multi-word denied phrase deliberately straddling the boundary) is
documented as an accepted limitation, since fixing it would require an
overlap window reconciled against masked output with no documented
length-preservation guarantee from AWS.
This commit is contained in:
spencer-burridge 2026-07-29 08:27:21 -05:00
parent 0d3f65439d
commit 8a63254fe2
2 changed files with 351 additions and 22 deletions

View file

@ -30,6 +30,7 @@ from typing import (
import copy
from collections.abc import Mapping
from datetime import datetime, timezone
from functools import reduce
import httpx
from fastapi import HTTPException
@ -83,6 +84,15 @@ from litellm.types.utils import (
)
GUARDRAIL_NAME = "bedrock"
# KNOWN LIMITATION (chunking, below): splitting an oversized message's text on
# a whitespace boundary (see `_nearest_whitespace_split_index`) prevents
# accidentally severing a single token -- one denied word, one PII pattern --
# across a chunk boundary. It does not stop a multi-word denied phrase
# deliberately positioned to straddle that boundary, since each fragment can
# scan clean independently. AWS's own guidance for this API acknowledges the
# same gap for input chunking with no documented resolution; closing it would
# require an overlap window reconciled against masked output, which AWS does
# not guarantee to be length-preserving. Accepted as out of scope.
_BEDROCK_DYNAMIC_BODY_DENYLIST = frozenset({"content", "source"})
# ApplyGuardrail's per-request "maximum input size in text units" quota is
# region/account/policy-dependent and cannot be predicted from config, so it is
@ -98,6 +108,18 @@ _BEDROCK_TOO_LARGE_ERROR_SUBSTRINGS = (
"too large",
"exceeds the maximum",
)
# Conservative starting guess for how much content (by character count) to send
# in one ApplyGuardrail call, used to pre-bin-pack content instead of always
# starting from the whole payload. This is NOT a correctness dependency: it only
# sets how many calls the common case takes. Any bin AWS still rejects as too
# large (because the real per-request text-unit cap for this account/region/
# policy is lower than this guess -- that cap is not knowable ahead of time and
# is not a fixed character count) falls back to the recursive bisection below,
# which self-corrects regardless of how wrong this guess was. So a too-generous
# guess here costs the same one extra probe-and-bisect round trip that pure
# reactive bisection would have paid anyway, while a well-tuned guess makes the
# common case a single pass instead of O(log n) round trips per request.
_BEDROCK_APPLY_GUARDRAIL_CHUNK_BUDGET_CHARS = 20_000
# 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.
@ -831,19 +853,28 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
# 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)
batches = (
self._bin_pack_bedrock_content(content, budget=_BEDROCK_APPLY_GUARDRAIL_CHUNK_BUDGET_CHARS)
if allow_chunking
else [content]
)
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,
)
responses = [
result
for batch in batches
for result in await self._apply_guardrail_content_with_chunking(
content=batch,
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
@ -1137,6 +1168,40 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
return True
return False
@staticmethod
def _bin_pack_bedrock_content(
content: list[BedrockContentItem],
budget: int,
) -> list[list[BedrockContentItem]]:
"""Pack whole content items, in order, into batches whose combined text
length stays within `budget`.
This is the fast-path half of the hybrid chunking strategy: bin-packing
at a conservative fixed budget keeps the common case at O(n / budget)
ApplyGuardrail calls instead of the O(log n) round trips pure reactive
bisection pays on every oversized request. An item whose own text
already exceeds `budget` is not split here -- it becomes its own
(still oversized) batch and is sent as-is; if AWS rejects that batch as
too large, `_apply_guardrail_content_with_chunking`'s existing
recursive-bisection fallback takes over for that batch only.
"""
if not content:
return [content]
def item_len(item: BedrockContentItem) -> int:
return len((item.get("text") or BedrockTextContent()).get("text") or "")
def add_item(
batches: tuple[tuple[BedrockContentItem, ...], ...],
item: BedrockContentItem,
) -> tuple[tuple[BedrockContentItem, ...], ...]:
if batches and sum(item_len(existing) for existing in batches[-1]) + item_len(item) <= budget:
return batches[:-1] + (batches[-1] + (item,),)
return batches + ((item,),)
packed = reduce(add_item, content, ())
return [list(batch) for batch in packed]
@staticmethod
def _split_bedrock_content(
content: list[BedrockContentItem],
@ -1145,12 +1210,31 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
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.
instead (a list of length 1 has no items left to bisect, but one very
long message is still a single content item) -- at the whitespace
character nearest the midpoint rather than a raw character index, so
the cut never lands inside a word/token. This is a plain, lossless
cut with no overlap: concatenating the two fragments in order always
reproduces the original text exactly, so merging back at
``_merge_logical_unit_outputs`` needs no reconciliation step.
Known, accepted limitation: whitespace splitting only guards against
*accidentally* severing a single token (one denied word, one PII
pattern) across the cut. It does not, and cannot without an overlap
window, stop a *multi-word* denied phrase deliberately positioned to
straddle the boundary -- each fragment can scan clean on its own and
still reassemble into the flagged phrase. AWS's own guidance on this
API acknowledges the same gap for input chunking ("a critical piece of
text could span two (or more) chunks if not carefully divided") with
no documented resolution, and overlap-and-reconcile was evaluated and
rejected for this PR: AWS's masking output has no documented
length-preservation guarantee, so reconciling an overlap region against
masked text is not sound in general. Out of scope for this PR.
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.
"""
if len(content) > 1:
midpoint = max(1, len(content) // 2)
@ -1160,16 +1244,35 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
text = text_content.get("text") or ""
if len(text) < 2:
return None
midpoint = len(text) // 2
split_at = BedrockGuardrail._nearest_whitespace_split_index(text)
qualifiers = text_content.get("qualifiers")
if qualifiers:
first_text = BedrockTextContent(text=text[:midpoint], qualifiers=qualifiers)
second_text = BedrockTextContent(text=text[midpoint:], qualifiers=qualifiers)
first_text = BedrockTextContent(text=text[:split_at], qualifiers=qualifiers)
second_text = BedrockTextContent(text=text[split_at:], qualifiers=qualifiers)
else:
first_text = BedrockTextContent(text=text[:midpoint])
second_text = BedrockTextContent(text=text[midpoint:])
first_text = BedrockTextContent(text=text[:split_at])
second_text = BedrockTextContent(text=text[split_at:])
return [BedrockContentItem(text=first_text)], [BedrockContentItem(text=second_text)]
@staticmethod
def _nearest_whitespace_split_index(text: str) -> int:
"""Return the index nearest `text`'s midpoint that falls on a
whitespace boundary, so splitting `text[:i]` / `text[i:]` there never
severs a word. Falls back to the raw midpoint when `text` has no
whitespace at all (a single giant token) -- still a correct, lossless
split, just no longer guaranteed word-safe for that pathological case.
"""
midpoint = len(text) // 2
left = text.rfind(" ", 0, midpoint)
right = text.find(" ", midpoint)
if left == -1 and right == -1:
return midpoint
if left == -1:
return right + 1
if right == -1:
return left + 1
return left + 1 if midpoint - left <= right - midpoint else right + 1
@staticmethod
def _is_input_too_large_validation_error(detail: object) -> bool:
"""True if `detail` is the AWS ValidationException message for input

View file

@ -16,11 +16,16 @@ import litellm
from litellm.caching.caching import DualCache
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
_BEDROCK_APPLY_GUARDRAIL_CHUNK_BUDGET_CHARS,
BedrockGuardrail,
_redact_pii_matches,
)
from litellm.proxy.utils import ProxyLogging
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
BedrockContentItem,
BedrockTextContent,
)
from litellm.types.utils import ModelResponse
@ -3898,3 +3903,224 @@ async def test_apply_guardrail_chunk_merge_preserves_masking_position():
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]"
@pytest.mark.asyncio
async def test_apply_guardrail_bin_packs_under_budget_content_with_no_probe_call():
"""Content that fits under the fixed chunk budget in one pre-packed batch
must be sent in exactly one ApplyGuardrail call -- no initial too-large
probe call, unlike pure reactive bisection which always pays that extra
round trip. Regression for: falling back to always trying the whole
unpacked content list first, rather than bin-packing before the first
attempt."""
guardrail = _bedrock_guardrail_for_chunk_tests()
# Five items, individually tiny, whose combined length exceeds the budget
# only when summed -- proves this triggers packing into multiple batches
# by SIZE, not by falling back to per-item chunking.
item_text = "x" * (_BEDROCK_APPLY_GUARDRAIL_CHUNK_BUDGET_CHARS // 2)
messages = [{"role": "user", "content": item_text} for _ in range(3)]
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
return _passing_bedrock_httpx_response(f"batch-{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
result = await guardrail.make_bedrock_api_request(
source="INPUT",
messages=messages,
request_data={"model": "bedrock-nova-micro"},
)
# Three items at budget/2 each pack two-per-batch (2 + 1), never a single
# oversized call and never a wasted whole-content probe: exactly 2 calls.
assert call_count == 2
assert result.get("action") == "NONE"
@pytest.mark.asyncio
async def test_apply_guardrail_small_content_makes_exactly_one_call():
"""Content that fits entirely within the budget in a single batch must
make exactly one ApplyGuardrail call -- confirms bin-packing does not
introduce an extra probe call for the common (small-request) case."""
guardrail = _bedrock_guardrail_for_chunk_tests()
messages = [
{"role": "user", "content": "short message one"},
{"role": "user", "content": "short message two"},
]
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()),
):
mock_post.return_value = _passing_bedrock_httpx_response("single-batch")
result = await guardrail.make_bedrock_api_request(
source="INPUT",
messages=messages,
request_data={"model": "bedrock-nova-micro"},
)
mock_post.assert_awaited_once()
assert result.get("action") == "NONE"
@pytest.mark.asyncio
async def test_apply_guardrail_batch_under_budget_still_rejected_falls_back_to_bisection():
"""A pre-packed batch that fits the fixed budget guess but is still
rejected by AWS as too large (a lower real per-account/region/policy cap)
must fall back to bisection for that batch only -- and any other batch
from the same request that AWS already accepted must not be re-sent."""
guardrail = _bedrock_guardrail_for_chunk_tests()
item_text = "x" * (_BEDROCK_APPLY_GUARDRAIL_CHUNK_BUDGET_CHARS // 2)
messages = [
{"role": "user", "content": item_text},
{"role": "user", "content": item_text},
{"role": "user", "content": item_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:
# First pre-packed batch (items 1+2): accepted immediately.
return _passing_bedrock_httpx_response("batch-1")
if call_count == 2:
# Second pre-packed batch (item 3 alone): rejected as too large
# despite fitting the fixed budget guess -- simulates a lower
# real-world per-account cap.
return _too_large_validation_httpx_response()
return _passing_bedrock_httpx_response(f"batch-2-bisected-{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
result = await guardrail.make_bedrock_api_request(
source="INPUT",
messages=messages,
request_data={"model": "bedrock-nova-micro"},
)
# batch-1 (1 call, accepted) + batch-2 (1 rejected + 2 bisected halves) = 4.
assert call_count == 4
assert result.get("action") == "NONE"
output_texts = [o.get("text") for o in result.get("outputs") or []]
# batch-2's two bisected text fragments came from the same original
# content item, so they are merged back into one combined output entry.
assert output_texts == ["batch-1", "batch-2-bisected-3batch-2-bisected-4"]
def test_split_bedrock_content_single_item_splits_on_whitespace_not_mid_word():
"""A single content item whose raw character midpoint would fall inside a
word must instead split at the nearest whitespace, so neither fragment
ends or begins mid-token. Regression for the Veria AI review finding: a
denied word/PII pattern straddling a raw character-midpoint cut could be
truncated on both fragments and scan clean on each, then reassemble into
the original unmasked text -- a detection bypass."""
# 20 'a's + space + 30 'b's: the raw character midpoint (25) falls inside
# the run of 'b's, proving the split must move off it to the nearest space.
text = ("a" * 20) + " " + ("b" * 30)
raw_midpoint = len(text) // 2
assert text[raw_midpoint] == "b"
content = [BedrockContentItem(text=BedrockTextContent(text=text))]
split_content = BedrockGuardrail._split_bedrock_content(content)
assert split_content is not None
first_half, second_half = split_content
first_text = first_half[0]["text"]["text"]
second_text = second_half[0]["text"]["text"]
# Lossless: concatenating the two fragments reproduces the original exactly.
assert first_text + second_text == text
# Word-safe: the split lands exactly on the whitespace boundary, not
# inside either the "a" or "b" run.
assert first_text == ("a" * 20) + " "
assert second_text == "b" * 30
def test_split_bedrock_content_single_item_with_no_whitespace_falls_back_to_midpoint():
"""A single giant token with no whitespace anywhere has no safe split
point, so the split must fall back to the raw character midpoint rather
than failing or looping."""
text = "a" * 40
content = [BedrockContentItem(text=BedrockTextContent(text=text))]
split_content = BedrockGuardrail._split_bedrock_content(content)
assert split_content is not None
first_half, second_half = split_content
first_text = first_half[0]["text"]["text"]
second_text = second_half[0]["text"]["text"]
assert first_text + second_text == text
assert len(first_text) == 20
assert len(second_text) == 20
def test_bin_pack_bedrock_content_packs_minimal_batches_within_budget():
"""Many medium items should pack into the minimal number of in-order
batches that each stay within budget, not one batch per item."""
items = [BedrockContentItem(text=BedrockTextContent(text="x" * 30)) for _ in range(10)]
batches = BedrockGuardrail._bin_pack_bedrock_content(items, budget=100)
assert sum(len(batch) for batch in batches) == 10
for batch in batches:
combined_len = sum(len(item["text"]["text"]) for item in batch)
assert combined_len <= 100
# 10 items * 30 chars = 300 chars at a 100-char budget packs into 3 batches
# of 3 items (90 chars) plus 1 batch of 1 item -- never one batch per item.
assert len(batches) == 4
def test_bin_pack_bedrock_content_oversized_single_item_becomes_its_own_batch():
"""An item whose own text already exceeds the budget must not be
pre-split here -- it becomes its own oversized batch, and only the
reactive bisection fallback (on an AWS rejection) may split it later."""
small_item = BedrockContentItem(text=BedrockTextContent(text="short"))
oversized_item = BedrockContentItem(text=BedrockTextContent(text="x" * 200))
items = [small_item, oversized_item, small_item]
batches = BedrockGuardrail._bin_pack_bedrock_content(items, budget=100)
assert batches == [[small_item], [oversized_item], [small_item]]
def test_bin_pack_bedrock_content_empty_content_makes_exactly_one_empty_batch():
"""Empty content must still pack into exactly one (empty) batch, matching
pre-bin-packing behavior of sending the content list as-is in one call --
bin-packing must not turn an empty request into zero ApplyGuardrail calls."""
assert BedrockGuardrail._bin_pack_bedrock_content([], budget=100) == [[]]