feat(bedrock): honor streaming buffer/sampling config for unbuffered post_call scans

This commit is contained in:
mateo-berri 2026-08-28 17:35:15 -07:00
parent d42c71b7ff
commit a90ad5fe5c
4 changed files with 253 additions and 1 deletions

View file

@ -52,7 +52,12 @@ from litellm.proxy.guardrails.anthropic_sse import (
model_response_text,
)
from litellm.secret_managers.main import get_secret_str
from litellm.types.guardrails import BedrockChecksConfigModel, GuardrailEventHooks
from litellm.types.guardrails import (
BedrockChecksConfigModel,
BedrockGuardrailStreamingParams,
GuardrailEventHooks,
LitellmParams,
)
from litellm.types.llms.openai import AllMessageValues, ChatCompletionUserMessage
from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
BedrockChecksMessage,
@ -221,9 +226,21 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
prompt_attack_threshold: float | None = 0.5,
pii_confidence_threshold: float | None = 0.5,
chunk_budget_chars: int = BEDROCK_APPLY_GUARDRAIL_CHUNK_BUDGET_CHARS,
streaming_buffer_until_moderated: bool | None = None,
streaming_sampling_rate: int | None = None,
streaming_end_of_stream_only: bool | None = None,
**kwargs,
):
self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
self._set_streaming_params(
BedrockGuardrailStreamingParams.from_extras(
{
"streaming_buffer_until_moderated": streaming_buffer_until_moderated,
"streaming_sampling_rate": streaming_sampling_rate,
"streaming_end_of_stream_only": streaming_end_of_stream_only,
}
)
)
self.guardrailIdentifier = guardrailIdentifier
self.guardrailVersion = guardrailVersion
self.guardrail_provider = "bedrock"
@ -278,6 +295,18 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
list(self.checks.keys()) if self.checks else None,
)
def _set_streaming_params(self, streaming_params: BedrockGuardrailStreamingParams) -> None:
self.streaming_buffer_until_moderated = streaming_params.streaming_buffer_until_moderated
self.streaming_sampling_rate = streaming_params.streaming_sampling_rate
self.streaming_end_of_stream_only = streaming_params.streaming_end_of_stream_only
def update_in_memory_litellm_params(self, litellm_params: LitellmParams) -> None:
super().update_in_memory_litellm_params(litellm_params)
self._set_streaming_params(BedrockGuardrailStreamingParams.from_extras(litellm_params.model_extra))
def _streams_incrementally(self) -> bool:
return not self.streaming_buffer_until_moderated and not self.mask_response_content
@classmethod
def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
return [
@ -2660,6 +2689,21 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
Collect content from the stream and run the bedrock OUTPUT scan
(post_call only validates the response).
"""
if self._streams_incrementally():
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
UnifiedLLMGuardrails,
)
async for streamed_chunk in UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook(
user_api_key_dict=user_api_key_dict,
response=response,
request_data=request_data,
guardrail_to_apply=self,
buffer_until_moderated_default=False,
):
yield streamed_chunk
return
# Import here to avoid circular imports
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
from litellm.main import stream_chunk_builder

View file

@ -11,6 +11,7 @@ def initialize_bedrock(litellm_params: LitellmParams, guardrail: Guardrail):
BedrockGuardrail,
)
streaming_params: Final = BedrockGuardrailStreamingParams.from_extras(litellm_params.model_extra)
_bedrock_callback: Final = BedrockGuardrail(
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
@ -38,6 +39,9 @@ def initialize_bedrock(litellm_params: LitellmParams, guardrail: Guardrail):
aws_bedrock_runtime_endpoint=litellm_params.aws_bedrock_runtime_endpoint,
experimental_use_latest_role_message_only=litellm_params.experimental_use_latest_role_message_only,
only_scan_new_messages=litellm_params.only_scan_new_messages or False,
streaming_buffer_until_moderated=streaming_params.streaming_buffer_until_moderated,
streaming_sampling_rate=streaming_params.streaming_sampling_rate,
streaming_end_of_stream_only=streaming_params.streaming_end_of_stream_only,
)
litellm.logging_callback_manager.add_litellm_callback(_bedrock_callback)
return _bedrock_callback

View file

@ -1,3 +1,4 @@
from collections.abc import Mapping
from datetime import datetime
from enum import Enum
from typing import Any, Final, Literal
@ -550,6 +551,37 @@ class BedrockGuardrailConfigModel(BaseModel):
)
class BedrockGuardrailStreamingParams(BaseModel):
streaming_buffer_until_moderated: bool = Field(
default=True,
description="If True (default), withhold every streamed chunk until the end-of-stream "
"ApplyGuardrail scan passes, so no flagged content reaches the client before a block. "
"If False, chunks stream through unbuffered, so flagged content can reach the client "
"before the scan finishes; a flagged scan still ends the stream, with a block message "
"when disable_exception_on_block is true and an in-stream error frame otherwise.",
)
streaming_sampling_rate: int = Field(
default=5,
ge=1,
description="When not buffering and not end-of-stream-only, scan the accumulated response "
"every Nth streamed chunk. Each sampled scan is a full ApplyGuardrail call that delays "
"that chunk, so lower values add latency and AWS text-unit cost.",
)
streaming_end_of_stream_only: bool = Field(
default=False,
description="When not buffering, skip per-chunk sampling and run one ApplyGuardrail scan "
"on the assembled response at end of stream. Combined with "
"streaming_buffer_until_moderated=false the full response streams live before the scan "
"and the scan result lands in guardrail_information; a flagged response still ends the "
"stream with a block message (disable_exception_on_block=true) or an error frame.",
)
@classmethod
def from_extras(cls, extras: Mapping[str, object] | None) -> "BedrockGuardrailStreamingParams":
source: Final[Mapping[str, object]] = extras or {}
return cls.model_validate({name: source[name] for name in cls.model_fields if source.get(name) is not None})
class LakeraV2GuardrailConfigModel(BaseModel):
"""Configuration parameters for the Lakera AI v2 guardrail"""

View file

@ -5345,3 +5345,175 @@ def test_initialize_bedrock_forwards_aws_external_id():
assert guardrail.optional_params["aws_external_id"] == "external-id-123"
finally:
litellm.logging_callback_manager.remove_callback_from_list_by_object(litellm.callbacks, guardrail)
def _chat_chunk(content: str, finish_reason: str | None) -> litellm.ModelResponseStream:
return litellm.ModelResponseStream(
id="tid",
choices=[
litellm.types.utils.StreamingChoices(
delta=litellm.types.utils.Delta(content=content, role="assistant"),
finish_reason=finish_reason,
index=0,
)
],
created=1,
model="gpt-4o-mini",
object="chat.completion.chunk",
)
def _streaming_litellm_params(**extras):
from litellm.types.guardrails import LitellmParams
return LitellmParams(
guardrail="bedrock",
mode="post_call",
guardrailIdentifier="test-id",
guardrailVersion="DRAFT",
**extras,
)
def test_initialize_bedrock_wires_streaming_flags():
from litellm.proxy.guardrails.guardrail_initializers import initialize_bedrock
configured = initialize_bedrock(
_streaming_litellm_params(
streaming_buffer_until_moderated=False,
streaming_sampling_rate=3,
streaming_end_of_stream_only=True,
),
{"guardrail_name": "bedrock-streaming"},
)
defaulted = initialize_bedrock(
_streaming_litellm_params(),
{"guardrail_name": "bedrock-defaults"},
)
for registered in (configured, defaulted):
litellm.logging_callback_manager.remove_callback_from_list_by_object(litellm.callbacks, registered)
assert configured.streaming_buffer_until_moderated is False
assert configured.streaming_sampling_rate == 3
assert configured.streaming_end_of_stream_only is True
assert defaulted.streaming_buffer_until_moderated is True
assert defaulted.streaming_sampling_rate == 5
assert defaulted.streaming_end_of_stream_only is False
def test_initialize_bedrock_rejects_non_positive_sampling_rate():
from pydantic import ValidationError
from litellm.proxy.guardrails.guardrail_initializers import initialize_bedrock
with pytest.raises(ValidationError):
initialize_bedrock(
_streaming_litellm_params(streaming_sampling_rate=0),
{"guardrail_name": "bedrock-bad-rate"},
)
def test_update_in_memory_litellm_params_round_trips_streaming_flags():
guardrail = BedrockGuardrail(
guardrail_name="bedrock-update",
guardrailIdentifier="test-id",
guardrailVersion="DRAFT",
)
guardrail.update_in_memory_litellm_params(
_streaming_litellm_params(
streaming_buffer_until_moderated=False,
streaming_sampling_rate=7,
streaming_end_of_stream_only=True,
)
)
assert guardrail.streaming_buffer_until_moderated is False
assert guardrail.streaming_sampling_rate == 7
assert guardrail.streaming_end_of_stream_only is True
guardrail.update_in_memory_litellm_params(_streaming_litellm_params())
assert guardrail.streaming_buffer_until_moderated is True
assert guardrail.streaming_sampling_rate == 5
assert guardrail.streaming_end_of_stream_only is False
async def _run_streaming_hook_recording_order(guardrail: BedrockGuardrail) -> list:
events = []
minimal = {"action": "NONE", "assessments": [], "outputs": []}
async def record_scan(*args, **kwargs):
events.append("scan")
return minimal
async def mock_stream():
yield _chat_chunk("Hello", None)
yield _chat_chunk(" world", None)
yield _chat_chunk("", "stop")
with patch.object(guardrail, "make_bedrock_api_request", AsyncMock(side_effect=record_scan)):
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=UserAPIKeyAuth(),
response=mock_stream(),
request_data={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hi"}]},
):
content = chunk.choices[0].delta.content if chunk.choices else None
events.append(("chunk", content))
return events
@pytest.mark.asyncio
async def test_unbuffered_end_of_stream_hook_yields_chunks_before_scan():
guardrail = BedrockGuardrail(
guardrail_name="bedrock-audit-mode",
guardrailIdentifier="test-id",
guardrailVersion="DRAFT",
event_hook=GuardrailEventHooks.post_call,
default_on=True,
streaming_buffer_until_moderated=False,
streaming_end_of_stream_only=True,
)
events = await _run_streaming_hook_recording_order(guardrail)
scan_index = events.index("scan")
chunk_events = [e for e in events if e != "scan"]
assert events.count("scan") == 1
assert [e for e in events[:scan_index] if e != "scan"] == chunk_events[: scan_index]
assert ("chunk", "Hello") in events[:scan_index]
assert ("chunk", " world") in events[:scan_index]
assert len(chunk_events) == 3
@pytest.mark.asyncio
async def test_buffered_default_hook_scans_before_any_chunk():
guardrail = BedrockGuardrail(
guardrail_name="bedrock-buffered-default",
guardrailIdentifier="test-id",
guardrailVersion="DRAFT",
event_hook=GuardrailEventHooks.post_call,
default_on=True,
)
events = await _run_streaming_hook_recording_order(guardrail)
assert events[0] == "scan"
assert all(e == "scan" or e[0] == "chunk" for e in events)
assert len([e for e in events if e != "scan"]) >= 1
@pytest.mark.asyncio
async def test_masking_keeps_buffered_path_even_when_unbuffered_configured():
guardrail = BedrockGuardrail(
guardrail_name="bedrock-mask-buffered",
guardrailIdentifier="test-id",
guardrailVersion="DRAFT",
event_hook=GuardrailEventHooks.post_call,
default_on=True,
mask_response_content=True,
streaming_buffer_until_moderated=False,
streaming_end_of_stream_only=True,
)
assert guardrail._streams_incrementally() is False
events = await _run_streaming_hook_recording_order(guardrail)
assert events[0] == "scan"