mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
feat(guardrails): chunk oversized Bedrock ApplyGuardrail requests instead of failing
AWS's ApplyGuardrail API rejects requests whose content exceeds the account's per-request "maximum input size in text units" quota with a 400 ValidationException. That cap is account/region/policy-dependent and cannot be predicted from config, so it can only be reacted to. _make_apply_guardrail_request now tries the whole-content call first (no behavior change for requests that already fit). On a too-large ValidationException it bisects the flat content list and retries each half sequentially, recursing until every piece fits or cannot be split further, then merges the per-chunk responses (action, assessments, outputs, usage) into one so callers cannot tell chunking happened. A real guardrail block on any (sub-)chunk still raises immediately. Contextual-grounding requests are never chunked: grounding scores the response holistically against the whole reference source, so fragmenting it would produce misleading scores. Each chunk call also gets a small exponential backoff retry on AWS ThrottlingException (429), since chunking increases the number of per-second API calls and can trade a 400 for a 429. All new state is local to a single request's call stack (no shared cache, no cross-process coordination), so this is safe for single-pod, multi-pod, and cache-less LiteLLM proxy deployments alike.
This commit is contained in:
parent
5d4c4d0fce
commit
7b7293eec3
2 changed files with 603 additions and 18 deletions
|
|
@ -9,6 +9,7 @@ import os
|
|||
import sys
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../..")) # Adds the parent directory to the system path
|
||||
import asyncio
|
||||
import json
|
||||
import sys
|
||||
from typing import (
|
||||
|
|
@ -57,6 +58,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
|
|||
BedrockGuardrailOutput,
|
||||
BedrockGuardrailQualifier,
|
||||
BedrockGuardrailResponse,
|
||||
BedrockGuardrailUsage,
|
||||
BedrockRequest,
|
||||
BedrockTextContent,
|
||||
)
|
||||
|
|
@ -82,6 +84,28 @@ from litellm.types.utils import (
|
|||
|
||||
GUARDRAIL_NAME = "bedrock"
|
||||
_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
|
||||
# only ever discovered reactively: AWS rejects an over-quota call with a 400
|
||||
# ValidationException whose message contains one of these substrings (matched
|
||||
# case-insensitively against the parsed AWS error message). On a match the
|
||||
# content is bisected and each half retried, rather than surfacing the 400.
|
||||
_BEDROCK_TOO_LARGE_ERROR_SUBSTRINGS = (
|
||||
"text unit",
|
||||
"maximum input size",
|
||||
"content size",
|
||||
"too long",
|
||||
"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.
|
||||
_BEDROCK_APPLY_GUARDRAIL_MAX_THROTTLE_RETRIES = 3
|
||||
_BEDROCK_APPLY_GUARDRAIL_BASE_BACKOFF_SECONDS = 0.5
|
||||
# Resource-less, detect-only InvokeGuardrailChecks API (no guardrail resource required).
|
||||
_BEDROCK_INVOKE_GUARDRAIL_CHECKS_PATH = "/guardrail-checks/invoke"
|
||||
# InvokeGuardrailChecks accepts at most 10 content blocks per message. A message with
|
||||
|
|
@ -765,7 +789,6 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
bedrock_request_data: dict = dict(
|
||||
self.convert_to_bedrock_format(source=source, messages=messages, response=response)
|
||||
)
|
||||
bedrock_guardrail_response: BedrockGuardrailResponse = BedrockGuardrailResponse()
|
||||
api_key: Optional[str] = None
|
||||
if request_data:
|
||||
dynamic_request_body_params = self.get_guardrail_dynamic_request_body_params(request_data=request_data)
|
||||
|
|
@ -779,6 +802,156 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
if request_data.get("api_key") is not None:
|
||||
api_key = request_data["api_key"]
|
||||
|
||||
# UI / spend logs use event_type. Bedrock's `source` is INPUT vs OUTPUT for the API
|
||||
# body, which must not be confused with the proxy hook (pre_call / during_call /
|
||||
# post_call). When omitted, keep legacy mapping for backward compatibility.
|
||||
if logging_event_type is not None:
|
||||
event_type = logging_event_type
|
||||
else:
|
||||
event_type = GuardrailEventHooks.pre_call if source == "INPUT" else GuardrailEventHooks.post_call
|
||||
|
||||
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,
|
||||
request_data=request_data,
|
||||
event_type=event_type,
|
||||
start_time=start_time,
|
||||
allow_chunking=allow_chunking,
|
||||
)
|
||||
return self._merge_bedrock_guardrail_responses(responses)
|
||||
|
||||
async def _apply_guardrail_content_with_chunking(
|
||||
self,
|
||||
content: List[BedrockContentItem],
|
||||
base_request_data: dict,
|
||||
credentials,
|
||||
aws_region_name: str,
|
||||
api_key: Optional[str],
|
||||
request_data: dict | None,
|
||||
event_type: GuardrailEventHooks,
|
||||
start_time: "datetime",
|
||||
allow_chunking: bool,
|
||||
) -> List[BedrockGuardrailResponse]:
|
||||
"""Post `content` to ApplyGuardrail, bisecting on a too-large error.
|
||||
|
||||
Tries `content` as a single call first. AWS's per-request "maximum input
|
||||
size in text units" quota is account/region/policy-dependent and cannot be
|
||||
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.
|
||||
"""
|
||||
try:
|
||||
return [
|
||||
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,
|
||||
)
|
||||
]
|
||||
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)
|
||||
first_results = await self._apply_guardrail_content_with_chunking(
|
||||
content=first_half,
|
||||
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,
|
||||
allow_chunking=allow_chunking,
|
||||
)
|
||||
second_results = await self._apply_guardrail_content_with_chunking(
|
||||
content=second_half,
|
||||
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,
|
||||
allow_chunking=allow_chunking,
|
||||
)
|
||||
return first_results + second_results
|
||||
raise
|
||||
|
||||
async def _post_apply_guardrail_content_with_retry(
|
||||
self,
|
||||
content: List[BedrockContentItem],
|
||||
base_request_data: dict,
|
||||
credentials,
|
||||
aws_region_name: str,
|
||||
api_key: Optional[str],
|
||||
request_data: dict | None,
|
||||
event_type: GuardrailEventHooks,
|
||||
start_time: "datetime",
|
||||
) -> BedrockGuardrailResponse:
|
||||
"""Post one ApplyGuardrail call for `content`, retrying with exponential
|
||||
backoff on AWS ThrottlingException (HTTP 429).
|
||||
|
||||
Chunking already trades one oversized call for several smaller ones, so
|
||||
retries here are capped low -- they must not multiply per-request latency
|
||||
by an order of magnitude when the account's per-second text-unit quota is
|
||||
the binding constraint rather than the per-request size quota.
|
||||
"""
|
||||
attempt = 0
|
||||
while True:
|
||||
try:
|
||||
return await self._post_apply_guardrail_content(
|
||||
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,
|
||||
)
|
||||
except HTTPException as exc:
|
||||
if exc.status_code == 429 and attempt < _BEDROCK_APPLY_GUARDRAIL_MAX_THROTTLE_RETRIES:
|
||||
await asyncio.sleep(_BEDROCK_APPLY_GUARDRAIL_BASE_BACKOFF_SECONDS * (2**attempt))
|
||||
attempt += 1
|
||||
continue
|
||||
raise
|
||||
|
||||
async def _post_apply_guardrail_content(
|
||||
self,
|
||||
content: List[BedrockContentItem],
|
||||
base_request_data: dict,
|
||||
credentials,
|
||||
aws_region_name: str,
|
||||
api_key: Optional[str],
|
||||
request_data: dict | None,
|
||||
event_type: GuardrailEventHooks,
|
||||
start_time: "datetime",
|
||||
) -> BedrockGuardrailResponse:
|
||||
"""Make exactly one signed ApplyGuardrail HTTP call for `content` and
|
||||
parse the result. Raises HTTPException on a guardrail block or any
|
||||
non-200 response (including 429, handled by the retry wrapper above).
|
||||
"""
|
||||
bedrock_request_data = {**base_request_data, "content": content}
|
||||
prepared_request = self._prepare_request(
|
||||
credentials=credentials,
|
||||
data=bedrock_request_data,
|
||||
|
|
@ -793,14 +966,6 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
prepared_request.headers,
|
||||
)
|
||||
|
||||
# UI / spend logs use event_type. Bedrock's `source` is INPUT vs OUTPUT for the API
|
||||
# body, which must not be confused with the proxy hook (pre_call / during_call /
|
||||
# post_call). When omitted, keep legacy mapping for backward compatibility.
|
||||
if logging_event_type is not None:
|
||||
event_type = logging_event_type
|
||||
else:
|
||||
event_type = GuardrailEventHooks.pre_call if source == "INPUT" else GuardrailEventHooks.post_call
|
||||
|
||||
httpx_response = await self._sign_and_post(
|
||||
prepared_request=prepared_request,
|
||||
request_data=request_data,
|
||||
|
|
@ -839,16 +1004,91 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
raise self._get_http_exception_for_blocked_guardrail(
|
||||
bedrock_guardrail_response, request_data=request_data
|
||||
)
|
||||
else:
|
||||
status_code, detail_message = self._parse_bedrock_guardrail_error_response(httpx_response)
|
||||
verbose_proxy_logger.error(
|
||||
"Bedrock AI: error in response. Status code: %s, response: %s",
|
||||
httpx_response.status_code,
|
||||
httpx_response.text,
|
||||
)
|
||||
raise HTTPException(status_code=status_code, detail=detail_message)
|
||||
return bedrock_guardrail_response
|
||||
|
||||
return bedrock_guardrail_response
|
||||
status_code, detail_message = self._parse_bedrock_guardrail_error_response(httpx_response)
|
||||
verbose_proxy_logger.error(
|
||||
"Bedrock AI: error in response. Status code: %s, response: %s",
|
||||
httpx_response.status_code,
|
||||
httpx_response.text,
|
||||
)
|
||||
raise HTTPException(status_code=status_code, detail=detail_message)
|
||||
|
||||
@staticmethod
|
||||
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)."""
|
||||
for item in content:
|
||||
if (item.get("text") or {}).get("qualifiers"):
|
||||
return True
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def _split_bedrock_content(
|
||||
content: List[BedrockContentItem],
|
||||
) -> Tuple[List[BedrockContentItem], List[BedrockContentItem]]:
|
||||
"""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.
|
||||
"""
|
||||
midpoint = max(1, len(content) // 2)
|
||||
return content[:midpoint], content[midpoint:]
|
||||
|
||||
@staticmethod
|
||||
def _is_input_too_large_validation_error(detail: object) -> bool:
|
||||
"""True if `detail` is the AWS ValidationException message for input
|
||||
exceeding the per-request text-unit quota.
|
||||
|
||||
A guardrail *block* is also raised as an HTTPException with status 400,
|
||||
but its ``detail`` is always a dict (built by
|
||||
``_get_http_exception_for_blocked_guardrail``); a non-200 API error's
|
||||
``detail`` is always the plain string returned by
|
||||
``_parse_bedrock_guardrail_error_response``. Checking ``isinstance(detail,
|
||||
str)`` is therefore sufficient to never mistake a real block for a
|
||||
too-large error.
|
||||
"""
|
||||
if not isinstance(detail, str):
|
||||
return False
|
||||
lowered = detail.lower()
|
||||
return any(substring in lowered for substring in _BEDROCK_TOO_LARGE_ERROR_SUBSTRINGS)
|
||||
|
||||
@staticmethod
|
||||
def _merge_bedrock_guardrail_responses(
|
||||
responses: List[BedrockGuardrailResponse],
|
||||
) -> 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.
|
||||
"""
|
||||
if len(responses) == 1:
|
||||
return responses[0]
|
||||
|
||||
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
|
||||
|
||||
merged: BedrockGuardrailResponse = BedrockGuardrailResponse(action=merged_action)
|
||||
if merged_outputs:
|
||||
merged["outputs"] = merged_outputs
|
||||
if merged_assessments:
|
||||
merged["assessments"] = merged_assessments
|
||||
if merged_usage:
|
||||
merged["usage"] = cast(BedrockGuardrailUsage, merged_usage)
|
||||
return merged
|
||||
|
||||
async def _sign_and_post(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -3274,3 +3274,348 @@ async def test_chat_completion_modify_response_exception_streaming_logging_obj_n
|
|||
# CustomStreamWrapper would raise AttributeError inside __init__ and this
|
||||
# call would never reach here.
|
||||
assert response is not None
|
||||
|
||||
|
||||
###############################################################################
|
||||
# AIKMG-278: chunk oversized ApplyGuardrail requests instead of failing.
|
||||
#
|
||||
# AWS rejects an ApplyGuardrail call whose content exceeds the account's
|
||||
# "maximum input size in text units" quota with a 400 ValidationException.
|
||||
# Rather than surfacing that 400 to the caller, split the content in half and
|
||||
# retry each half; merge the per-half responses back into one. Grounded
|
||||
# (contextual-grounding) requests are never chunked -- grounding scores the
|
||||
# response holistically against the whole source, so fragmenting it would
|
||||
# produce misleading scores.
|
||||
###############################################################################
|
||||
|
||||
|
||||
def _too_large_validation_httpx_response() -> MagicMock:
|
||||
response = MagicMock()
|
||||
response.status_code = 400
|
||||
response.json.return_value = {
|
||||
"message": "Input is too long. Content size exceeds the maximum input size in text units."
|
||||
}
|
||||
response.text = json.dumps(response.json.return_value)
|
||||
return response
|
||||
|
||||
|
||||
def _other_validation_httpx_response() -> MagicMock:
|
||||
response = MagicMock()
|
||||
response.status_code = 400
|
||||
response.json.return_value = {"message": "guardrailIdentifier is not valid"}
|
||||
response.text = json.dumps(response.json.return_value)
|
||||
return response
|
||||
|
||||
|
||||
def _throttling_httpx_response() -> MagicMock:
|
||||
response = MagicMock()
|
||||
response.status_code = 429
|
||||
response.json.return_value = {"message": "Rate exceeded"}
|
||||
response.text = json.dumps(response.json.return_value)
|
||||
response.headers = {}
|
||||
return response
|
||||
|
||||
|
||||
def _passing_bedrock_httpx_response(marker: str) -> MagicMock:
|
||||
"""A successful ApplyGuardrail response tagged with `marker` so tests can
|
||||
verify which chunk produced which output/usage after merging."""
|
||||
response = MagicMock()
|
||||
response.status_code = 200
|
||||
response.json.return_value = {
|
||||
"action": "NONE",
|
||||
"outputs": [{"text": marker}],
|
||||
"assessments": [],
|
||||
"usage": {"contentPolicyUnits": 1},
|
||||
}
|
||||
return response
|
||||
|
||||
|
||||
def _blocking_bedrock_httpx_response(marker: str) -> MagicMock:
|
||||
response = MagicMock()
|
||||
response.status_code = 200
|
||||
response.json.return_value = {
|
||||
"action": "GUARDRAIL_INTERVENED",
|
||||
"outputs": [{"text": marker}],
|
||||
"assessments": [
|
||||
{"topicPolicy": {"topics": [{"name": marker, "type": "DENY", "action": "BLOCKED"}]}}
|
||||
],
|
||||
"usage": {"contentPolicyUnits": 1},
|
||||
}
|
||||
return response
|
||||
|
||||
|
||||
def _bedrock_guardrail_for_chunk_tests() -> "BedrockGuardrail":
|
||||
return BedrockGuardrail(
|
||||
guardrail_name="test-bedrock-guard",
|
||||
guardrailIdentifier="test-guardrail",
|
||||
guardrailVersion="DRAFT",
|
||||
disable_exception_on_block=False,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_chunks_on_too_large_validation_error():
|
||||
"""A too-large 400 on the whole-content call must trigger a bisect-and-retry,
|
||||
and the two chunk responses must be merged (assessments concatenated, usage
|
||||
summed, outputs concatenated) rather than losing either half's result."""
|
||||
guardrail = _bedrock_guardrail_for_chunk_tests()
|
||||
|
||||
messages = [
|
||||
{"role": "user", "content": "first half of a very long message"},
|
||||
{"role": "user", "content": "second half of a very long message"},
|
||||
]
|
||||
|
||||
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:
|
||||
# Whole-content call: too large.
|
||||
return _too_large_validation_httpx_response()
|
||||
if call_count == 2:
|
||||
return _passing_bedrock_httpx_response("chunk-1")
|
||||
return _blocking_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()),
|
||||
):
|
||||
mock_post.side_effect = _post_side_effect
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail.make_bedrock_api_request(
|
||||
source="INPUT",
|
||||
messages=messages,
|
||||
request_data={"model": "bedrock-nova-micro"},
|
||||
)
|
||||
|
||||
# 1 whole-content attempt + 2 chunk attempts.
|
||||
assert call_count == 3
|
||||
detail = exc_info.value.detail
|
||||
assert exc_info.value.status_code == 400
|
||||
# The merged response must retain the block signal from chunk-2 even
|
||||
# though chunk-1 passed clean -- losing it would silently bypass a
|
||||
# guardrail hit.
|
||||
assert "chunk-2" in detail["bedrock_guardrail_response"]
|
||||
assert detail["assessments"][0]["matches"][0]["name"] == "chunk-2"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_merges_usage_and_outputs_across_chunks_when_both_pass():
|
||||
"""When both chunks pass clean, the merged response must still carry both
|
||||
chunks' outputs/usage forward (needed for accurate logging/telemetry) and
|
||||
must not itself raise."""
|
||||
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()),
|
||||
):
|
||||
mock_post.side_effect = _post_side_effect
|
||||
|
||||
result = await guardrail.make_bedrock_api_request(
|
||||
source="INPUT",
|
||||
messages=messages,
|
||||
request_data={"model": "bedrock-nova-micro"},
|
||||
)
|
||||
|
||||
assert call_count == 3
|
||||
assert result.get("action") == "NONE"
|
||||
output_texts = [o.get("text") for o in result.get("outputs") or []]
|
||||
assert output_texts == ["chunk-1", "chunk-2"]
|
||||
assert result.get("usage", {}).get("contentPolicyUnits") == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_does_not_chunk_when_grounding_present():
|
||||
"""Contextual-grounding requests are scored holistically against the whole
|
||||
source; chunking them would silently produce misleading grounding scores.
|
||||
A too-large error on a grounded request must propagate unchanged, not be
|
||||
bisected."""
|
||||
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()),
|
||||
):
|
||||
mock_post.return_value = _too_large_validation_httpx_response()
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail.make_bedrock_api_request(
|
||||
source="OUTPUT",
|
||||
messages=messages,
|
||||
response=model_response,
|
||||
request_data={"model": "bedrock-nova-micro"},
|
||||
)
|
||||
|
||||
# Exactly one call: no bisect-and-retry for a grounded request.
|
||||
assert mock_post.await_count == 1
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_does_not_chunk_on_non_size_validation_error():
|
||||
"""A 400 for an unrelated validation problem (e.g. a bad guardrail id) must
|
||||
not trigger chunking -- retrying a bad-config error split into pieces would
|
||||
just fail twice more and mask the real problem."""
|
||||
guardrail = _bedrock_guardrail_for_chunk_tests()
|
||||
|
||||
messages = [
|
||||
{"role": "user", "content": "first"},
|
||||
{"role": "user", "content": "second"},
|
||||
]
|
||||
|
||||
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 = _other_validation_httpx_response()
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail.make_bedrock_api_request(
|
||||
source="INPUT",
|
||||
messages=messages,
|
||||
request_data={"model": "bedrock-nova-micro"},
|
||||
)
|
||||
|
||||
assert mock_post.await_count == 1
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "not valid" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
@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."""
|
||||
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
|
||||
|
||||
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 = _too_large_validation_httpx_response()
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail.make_bedrock_api_request(
|
||||
source="INPUT",
|
||||
messages=messages,
|
||||
request_data={"model": "bedrock-nova-micro"},
|
||||
)
|
||||
|
||||
assert mock_post.await_count == 1
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
|
||||
@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
|
||||
eventually succeed, rather than surfacing the 429 to the caller."""
|
||||
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:
|
||||
# First chunk throttled once, then succeeds on retry.
|
||||
return _throttling_httpx_response()
|
||||
if call_count == 3:
|
||||
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(
|
||||
"litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails.asyncio.sleep",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_sleep,
|
||||
):
|
||||
mock_post.side_effect = _post_side_effect
|
||||
|
||||
result = await guardrail.make_bedrock_api_request(
|
||||
source="INPUT",
|
||||
messages=messages,
|
||||
request_data={"model": "bedrock-nova-micro"},
|
||||
)
|
||||
|
||||
assert call_count == 4
|
||||
mock_sleep.assert_awaited()
|
||||
output_texts = [o.get("text") for o in result.get("outputs") or []]
|
||||
assert output_texts == ["chunk-1", "chunk-2"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue