mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(bedrock): route streamed responses-API output through the unified guardrail
Streamed /v1/responses returned 500 whenever a Bedrock post_call guardrail was enabled: the hook fed responses-API events into stream_chunk_builder, which only understands chat-completions chunks, and the wrapped KeyError surfaced as litellm.APIError before any ApplyGuardrail scan ran. Delegate responses-API routes to UnifiedLLMGuardrails, whose translation layer scans the assembled response at end of stream and only then releases the buffered events, so flagged content never reaches the client.
This commit is contained in:
parent
4d7144160a
commit
7eb757a49b
2 changed files with 106 additions and 0 deletions
|
|
@ -30,6 +30,7 @@ from litellm.caching import DualCache
|
|||
from litellm.constants import BEDROCK_APPLY_GUARDRAIL_CHUNK_BUDGET_CHARS
|
||||
from litellm.exceptions import ModifyResponseException
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route
|
||||
from litellm.litellm_core_utils.core_helpers import redact_nested_match_and_regex_keys
|
||||
from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import bedrock_guardrail_cost
|
||||
from litellm.llms.anthropic.chat.guardrail_translation.handler import AnthropicMessagesHandler
|
||||
|
|
@ -206,6 +207,16 @@ def _redact_assessment_match_fields(assessments: list[dict]) -> list[dict]:
|
|||
return redacted if isinstance(redacted, list) else assessments
|
||||
|
||||
|
||||
_RESPONSES_API_CALL_TYPES: Final = frozenset({CallTypes.responses, CallTypes.aresponses})
|
||||
|
||||
|
||||
def _is_responses_api_route(request_route: str | None) -> bool:
|
||||
if request_route is None:
|
||||
return False
|
||||
call_types: Final = get_call_types_for_route(request_route)
|
||||
return call_types is not None and any(call_type in _RESPONSES_API_CALL_TYPES for call_type in call_types)
|
||||
|
||||
|
||||
class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
||||
# During-call must use async_moderation_hook (not unified apply_guardrail), otherwise
|
||||
# OpenAI translation always passes input_type="request" and spend/UI show PRE-CALL.
|
||||
|
|
@ -2660,6 +2671,24 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
Collect content from the stream and run the bedrock OUTPUT scan
|
||||
(post_call only validates the response).
|
||||
"""
|
||||
# Responses-API events are neither chat-completions chunks nor raw
|
||||
# Anthropic SSE, so the assembly below cannot scan them; the unified
|
||||
# guardrail's translation layer can, with buffering semantics kept.
|
||||
if _is_responses_api_route(user_api_key_dict.request_route):
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
|
||||
UnifiedLLMGuardrails,
|
||||
)
|
||||
|
||||
async for translated_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=True,
|
||||
):
|
||||
yield translated_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
|
||||
|
|
|
|||
|
|
@ -5345,3 +5345,80 @@ 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 _responses_stream_events() -> list:
|
||||
from litellm.types.llms.openai import (
|
||||
OutputTextDeltaEvent,
|
||||
ResponseCompletedEvent,
|
||||
ResponsesAPIResponse,
|
||||
ResponsesAPIStreamEvents,
|
||||
)
|
||||
|
||||
deltas = [
|
||||
OutputTextDeltaEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA,
|
||||
item_id="msg_lit6457",
|
||||
output_index=0,
|
||||
content_index=0,
|
||||
delta=part,
|
||||
)
|
||||
for part in ("Hello", " world")
|
||||
]
|
||||
completed = ResponseCompletedEvent(
|
||||
type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
|
||||
response=ResponsesAPIResponse(
|
||||
id="resp_lit6457",
|
||||
created_at=1234567890,
|
||||
model="gpt-4o",
|
||||
object="response",
|
||||
status="completed",
|
||||
output=[
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_lit6457",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "Hello world"}],
|
||||
}
|
||||
],
|
||||
),
|
||||
)
|
||||
return [*deltas, completed]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_api_stream_scans_output_and_replays_buffered_events():
|
||||
"""Streamed /v1/responses events must be scanned via the unified translation
|
||||
layer, not fed to stream_chunk_builder (which raises APIError on them)."""
|
||||
guardrail = BedrockGuardrail(
|
||||
guardrail_name="bedrock-responses-stream",
|
||||
guardrailIdentifier="test-id",
|
||||
guardrailVersion="DRAFT",
|
||||
event_hook=GuardrailEventHooks.post_call,
|
||||
default_on=True,
|
||||
)
|
||||
stream_events = _responses_stream_events()
|
||||
order = []
|
||||
yielded = []
|
||||
|
||||
async def record_scan(*args, **kwargs):
|
||||
order.append("scan")
|
||||
return {"action": "NONE", "assessments": [], "outputs": []}
|
||||
|
||||
async def mock_stream():
|
||||
for event in stream_events:
|
||||
yield event
|
||||
|
||||
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(api_key="test-key", request_route="/v1/responses"),
|
||||
response=mock_stream(),
|
||||
request_data={"model": "gpt-4o", "input": "hi"},
|
||||
):
|
||||
order.append("chunk")
|
||||
yielded.append(chunk)
|
||||
|
||||
assert order == ["scan", "chunk", "chunk", "chunk"]
|
||||
assert len(yielded) == len(stream_events)
|
||||
assert all(emitted is original for emitted, original in zip(yielded, stream_events))
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue