From 576e41797ef4cb80821e3c2924f2f1e598b77621 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Thu, 2 Jul 2026 16:13:51 +0000 Subject: [PATCH] fix: handle Anthropic streaming guardrail blocks --- .../chat/guardrail_translation/handler.py | 46 +++++++++++++++---- .../base_llm/guardrail_translation/utils.py | 10 ++-- .../unified_guardrail/unified_guardrail.py | 29 ++++++++++-- .../test_anthropic_streaming_block.py | 34 +++++++++++++- 4 files changed, 99 insertions(+), 20 deletions(-) diff --git a/litellm/llms/anthropic/chat/guardrail_translation/handler.py b/litellm/llms/anthropic/chat/guardrail_translation/handler.py index 469acd9c8cd..7000c20d9c4 100644 --- a/litellm/llms/anthropic/chat/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/chat/guardrail_translation/handler.py @@ -73,6 +73,22 @@ class AnthropicMessagesHandler(BaseTranslation): super().__init__() self.adapter = LiteLLMAnthropicMessagesAdapter() + @staticmethod + def _build_streaming_usage_response( + responses_so_far: list[Any], + request_data: Optional[dict], + ) -> Optional[ModelResponse]: + chunks = tuple(response for response in responses_so_far if isinstance(response, (str, bytes))) + if not chunks: + return None + try: + return AnthropicPassthroughLoggingHandler._build_usage_only_response_from_chunks( + all_chunks=chunks, + model=str((request_data or {}).get("model") or ""), + ) + except (AttributeError, TypeError, ValueError): + return None + def build_block_sse_chunks( self, exc: "ModifyResponseException", @@ -557,6 +573,8 @@ class AnthropicMessagesHandler(BaseTranslation): Get the string so far, check the apply guardrail to the string so far, and return the list of responses so far. """ + from litellm.integrations.custom_guardrail import ModifyResponseException + has_ended = self._check_streaming_has_ended(responses_so_far) if has_ended: # build the model response from the responses_so_far @@ -581,25 +599,35 @@ class AnthropicMessagesHandler(BaseTranslation): if tool_calls_list: guardrail_inputs["tool_calls"] = tool_calls_list - _guardrailed_inputs = ( - await guardrail_to_apply.apply_guardrail( # allow rejecting the response, if invalid + try: + _guardrailed_inputs = await guardrail_to_apply.apply_guardrail( inputs=guardrail_inputs, request_data=request_data if request_data is not None else {}, input_type="response", logging_obj=litellm_logging_obj, ) - ) + except ModifyResponseException as e: + if e.original_response is None: + e.original_response = built_response or self._build_streaming_usage_response( + responses_so_far, request_data + ) + raise else: verbose_proxy_logger.debug("Skipping output guardrail - model response has no choices") return responses_so_far string_so_far = self.get_streaming_string_so_far(responses_so_far) - _guardrailed_inputs = await guardrail_to_apply.apply_guardrail( # allow rejecting the response, if invalid - inputs={"texts": [string_so_far]}, - request_data=request_data if request_data is not None else {}, - input_type="response", - logging_obj=litellm_logging_obj, - ) + try: + _guardrailed_inputs = await guardrail_to_apply.apply_guardrail( + inputs={"texts": [string_so_far]}, + request_data=request_data if request_data is not None else {}, + input_type="response", + logging_obj=litellm_logging_obj, + ) + except ModifyResponseException as e: + if e.original_response is None: + e.original_response = self._build_streaming_usage_response(responses_so_far, request_data) + raise return responses_so_far def _prepare_request_data( diff --git a/litellm/llms/base_llm/guardrail_translation/utils.py b/litellm/llms/base_llm/guardrail_translation/utils.py index aade7180617..02dd7867d69 100644 --- a/litellm/llms/base_llm/guardrail_translation/utils.py +++ b/litellm/llms/base_llm/guardrail_translation/utils.py @@ -22,14 +22,14 @@ def blocked_response_usage(original_response: Optional[Any]) -> AnthropicUsage: elif original_response is not None: usage_obj = getattr(original_response, "usage", None) - def _tokens(key: str) -> int: + def _tokens(key: str, fallback_key: str) -> int: if isinstance(usage_obj, dict): - return int(usage_obj.get(key, 0) or 0) - return int(getattr(usage_obj, key, 0) or 0) + return int(usage_obj.get(key, usage_obj.get(fallback_key, 0)) or 0) + return int(getattr(usage_obj, key, getattr(usage_obj, fallback_key, 0)) or 0) return AnthropicUsage( - input_tokens=_tokens("input_tokens"), - output_tokens=_tokens("output_tokens"), + input_tokens=_tokens("input_tokens", "prompt_tokens"), + output_tokens=_tokens("output_tokens", "completion_tokens"), ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index 531d9a0c412..e0b387e92c0 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -384,6 +384,8 @@ class UnifiedLLMGuardrails(CustomLogger): call_type = None chunk_counter = 0 responses_so_far: List[Any] = [] + responses_yielded: list[Any] = [] + pending_end_of_stream_items: list[Any] = [] # Whether any real response chunk has been forwarded to the client. # Drives how a block terminates the stream: continue the in-progress # message (True) vs emit a standalone block message (False, buffered). @@ -415,8 +417,16 @@ class UnifiedLLMGuardrails(CustomLogger): # moderation runs below. if end_of_stream_only: if not buffer_until_moderated: - chunks_yielded = True - yield item + endpoint_translation = endpoint_guardrail_translation_mappings[CallTypes(call_type)]() + stream_has_ended = hasattr( + endpoint_translation, "_check_streaming_has_ended" + ) and endpoint_translation._check_streaming_has_ended(responses_so_far) + if pending_end_of_stream_items or stream_has_ended: + pending_end_of_stream_items.append(item) + else: + chunks_yielded = True + responses_yielded.append(item) + yield item continue # Process chunk based on sampling rate @@ -447,6 +457,8 @@ class UnifiedLLMGuardrails(CustomLogger): request_data=request_data, ) except ModifyResponseException as e: + if e.original_response is None: + e.original_response = responses_so_far # Guardrail blocked the response mid-stream. Emit a clean # terminating SSE sequence delivering the block message # instead of letting the exception propagate into a bare @@ -460,7 +472,7 @@ class UnifiedLLMGuardrails(CustomLogger): e, endpoint_translation, stream_started=chunks_yielded, - responses_so_far=responses_so_far[:-1], + responses_so_far=responses_yielded, ): yield block_chunk return @@ -491,9 +503,11 @@ class UnifiedLLMGuardrails(CustomLogger): return raise chunks_yielded = True + responses_yielded.append(original_item) yield original_item else: chunks_yielded = True + responses_yielded.append(item) yield item # Stream has ended - do final processing with all collected chunks @@ -527,7 +541,12 @@ class UnifiedLLMGuardrails(CustomLogger): if buffered_items is not None: for buffered_item in buffered_items: yield buffered_item + for pending_item in pending_end_of_stream_items: + responses_yielded.append(pending_item) + yield pending_item except ModifyResponseException as e: + if e.original_response is None: + e.original_response = responses_so_far # Block detected during end-of-stream processing. Emit a clean # terminating SSE sequence with the block message rather than # propagating into a bare error blob that truncates the stream. @@ -535,8 +554,8 @@ class UnifiedLLMGuardrails(CustomLogger): async for block_chunk in self._handle_streaming_block( e, endpoint_translation, - stream_started=chunks_yielded, - responses_so_far=responses_so_far, + stream_started=bool(responses_yielded), + responses_so_far=responses_yielded, ): yield block_chunk return diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_anthropic_streaming_block.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_anthropic_streaming_block.py index 69792fb2449..48109a632eb 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_anthropic_streaming_block.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_anthropic_streaming_block.py @@ -117,12 +117,13 @@ def _parse_sse_event_types(raw: str) -> List[str]: return event_types -async def _run_hook(end: bool, sampling_rate: int = 1) -> str: +async def _run_hook(end: bool, sampling_rate: int = 1, end_of_stream_only: bool = False) -> str: guardrail = _BlockingGuardrail(guardrail_name="test-blocking-guardrail", event_hook="post_call") # sampling_rate controls how many chunks are forwarded before the block # fires: 1 blocks on the first chunk (nothing sent yet); >1 forwards earlier # chunks first, exercising the mid-stream "continue the message" path. guardrail.streaming_sampling_rate = sampling_rate + guardrail.streaming_end_of_stream_only = end_of_stream_only unified_guardrail = UnifiedLLMGuardrails() user_api_key_dict = UserAPIKeyAuth(api_key="test", request_route="/v1/messages") @@ -159,6 +160,21 @@ def _assert_clean_block_termination(raw: str) -> None: assert any('"stop_reason"' in block and "message_delta" in block for block in raw.split("\n\n")) +def _parse_sse_payloads(raw: str) -> List[dict]: + payloads = [] + for block in raw.split("\n\n"): + for line in block.strip().split("\n"): + if line.startswith("data:"): + payload = line[len("data:") :].strip() + try: + parsed = json.loads(payload) + except json.JSONDecodeError: + continue + if isinstance(parsed, dict): + payloads.append(parsed) + return payloads + + @pytest.mark.asyncio async def test_mid_stream_block_emits_clean_anthropic_sse(): """Per-chunk block: a clean SSE termination with the block message, no error blob.""" @@ -185,6 +201,22 @@ async def test_mid_stream_block_after_prior_chunks_continues_message(): _assert_clean_block_termination(raw) +@pytest.mark.asyncio +async def test_end_of_stream_only_block_does_not_append_after_message_stop(): + raw = await _run_hook(end=True, end_of_stream_only=True) + event_types = _parse_sse_event_types(raw) + message_delta_usages = [ + payload.get("usage", {}).get("output_tokens") + for payload in _parse_sse_payloads(raw) + if payload.get("type") == "message_delta" + ] + + assert BLOCK_MESSAGE in raw + assert event_types.count("message_stop") == 1 + assert event_types[-1] == "message_stop" + assert message_delta_usages[-1] == 5 + + class TestContentBlockState: """`_content_block_state` must reflect the true open/last block index across the two chunk formats the stream can carry (multi-event bytes, parsed dict),