From 321345d4c8a7ae4ec05d036fbc9f5b7e31029d86 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Fri, 3 Jul 2026 09:27:31 +0530 Subject: [PATCH] feat: litellm oss staging (#31935) * fix(prometheus): bound per-request budget metric emission with a timeout (#31632) * fix(prometheus): bound per-request budget metric emission with a timeout Wrap the per-request budget-metric gather in asyncio.wait_for so a slow Redis or DB lookup cannot consume the whole LoggingWorker watchdog and get the success-logging event cancelled. On timeout the emission is skipped in isolation; budget gauges are still refreshed by the periodic cron. The timeout is configurable via PROMETHEUS_BUDGET_METRICS_PER_REQUEST_TIMEOUT and defaults to 5.0 seconds, falling back to the default on an invalid value instead of raising * fix(prometheus): reject non-finite and non-positive budget-metrics timeout env float() accepts 0, negatives, nan and inf, which bypass the fallback: a value <= 0 makes asyncio.wait_for time out immediately and skip every per-request emission, and inf reintroduces the unbounded wait the timeout was meant to bound. Validate the parsed value is finite and greater than zero before using it, otherwise fall back to the default * fix: report the blocked LLM response's real token usage (#31217) When a guardrail blocks a post-call response, the synthetic violation response reported hard-coded zero usage, discarding the token usage the upstream call had already consumed. Fix the root cause rather than re-counting tokens: - Add an optional `original_response` field to ModifyResponseException. - The unified guardrail's post-call success hook attaches the blocked LLM response to the exception. - The /v1/messages and OpenAI-format (/v1/chat/completions, /v1/completions) block handlers report `original_response.usage` directly. Pre-call blocks never invoked the LLM, so usage is zero. Mock-based tests cover the helper (returns original usage / zero), the success hook attaching original_response, and the endpoint reporting it end-to-end. Co-authored-by: Claude Opus 4.8 (1M context) * feat(guardrails): buffer + cleanly terminate streamed responses on block (#31389) Streaming moderation improvements for the unified guardrail post-call streaming iterator hook: - streaming_buffer_until_moderated: withhold all chunks until end-of-stream moderation passes, then release the original response (clean) or only the block message (blocked) -- the original content is never delivered on a block. Snapshot chunks with a shallow list() copy (end-of-stream builds a separate assembled response; chunks aren't mutated in place). - Clean Anthropic SSE on block: synthesize a well-formed termination sequence instead of a bare data: {"error": ...} blob that truncates the stream. Provider-specific synthesis lives in AnthropicMessagesHandler via build_block_sse_chunks (format-agnostic routing stays in the hook). - Mid-stream blocks continue the in-progress message (close open content block, append block message, terminate) rather than emitting a second message_start, which clients reject. Standalone envelope only when no chunks were sent (buffered path). - ModifyResponseException imported under TYPE_CHECKING + locally at runtime to avoid a module-level cyclic import. Adds regression tests for buffering (content withheld on block) and mid-stream continuation (single message_start). Co-authored-by: Claude Opus 4.8 (1M context) * fix: report real usage on streaming blocks, disable buffered mode for content-rewriting guardrails - _standalone_block_chunks and _block_continuation_chunks now read real token usage from ModifyResponseException.original_response instead of hardcoding zero, matching the non-streaming _blocked_response_usage path. Shared helper moved to guardrail_translation/utils.py. - streaming_buffer_until_moderated is now forced off when the guardrail has mask_response_content=True, since buffered replay releases the withheld original chunks verbatim -- unsafe for a guardrail that rewrites content (e.g. PII masking). - Fix inverted streaming-flag precedence comment. * style: ruff format after greploop fixes * fix: handle Anthropic streaming guardrail blocks * fix(responses): check terminal event type for streaming guardrail end-of-stream detection _check_streaming_has_ended assumed responses_so_far held ModelResponse objects with .choices, but for the Responses API the accumulated chunks are raw SSE event dicts, causing an AttributeError on every call * fix: preserve Anthropic blocked stream usage --------- Co-authored-by: FERNANDO IZAR Co-authored-by: Joseph Barker <156112794+seph-barker@users.noreply.github.com> Co-authored-by: Claude Opus 4.8 (1M context) Co-authored-by: Cursor Agent --- litellm/exceptions.py | 6 + litellm/integrations/prometheus.py | 41 +- .../chat/guardrail_translation/handler.py | 199 +++++++++- .../guardrail_translation/base_translation.py | 29 +- .../base_llm/guardrail_translation/utils.py | 92 ++++- .../guardrail_translation/handler.py | 10 +- .../proxy/anthropic_endpoints/endpoints.py | 9 +- .../unified_guardrail/unified_guardrail.py | 193 ++++++++-- litellm/proxy/proxy_server.py | 33 +- .../test_prometheus_budget_metrics_timeout.py | 168 ++++++++ .../anthropic_endpoints/test_endpoints.py | 76 +++- .../test_anthropic_streaming_block.py | 362 ++++++++++++++++++ .../test_streaming_buffer_until_moderated.py | 173 +++++++++ .../proxy/test_blocked_response_usage.py | 84 ++++ 14 files changed, 1412 insertions(+), 63 deletions(-) create mode 100644 tests/test_litellm/integrations/test_prometheus_budget_metrics_timeout.py create mode 100644 tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_anthropic_streaming_block.py create mode 100644 tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_streaming_buffer_until_moderated.py create mode 100644 tests/test_litellm/proxy/test_blocked_response_usage.py diff --git a/litellm/exceptions.py b/litellm/exceptions.py index d97ba347b07..adf7b3ef05a 100644 --- a/litellm/exceptions.py +++ b/litellm/exceptions.py @@ -1165,12 +1165,18 @@ class ModifyResponseException(Exception): request_data: Dict[str, Any], guardrail_name: Optional[str] = None, detection_info: Optional[Dict[str, Any]] = None, + original_response: Optional[Any] = None, ): self.message = message self.model = model self.request_data = request_data self.guardrail_name = guardrail_name self.detection_info = detection_info or {} + # The LLM response that was blocked (post-call). Carries the real token + # usage the upstream call consumed, so the synthetic block response can + # report it instead of discarding it. None for pre-call blocks (the LLM + # was never invoked). + self.original_response = original_response super().__init__(message) diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 4ebe312e301..0f89e21abda 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -4,6 +4,7 @@ from __future__ import annotations import asyncio +import math import os import sys from datetime import datetime, timedelta @@ -65,6 +66,26 @@ if TYPE_CHECKING: else: AsyncIOScheduler = Any +_DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT = 5.0 + + +def _get_budget_metrics_per_request_timeout() -> float: + raw = os.getenv("PROMETHEUS_BUDGET_METRICS_PER_REQUEST_TIMEOUT") + if raw is None: + return _DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT + try: + parsed = float(raw) + except ValueError: + parsed = None + if parsed is None or not math.isfinite(parsed) or parsed <= 0: + verbose_logger.debug( + "[Non-Blocking] Prometheus: invalid PROMETHEUS_BUDGET_METRICS_PER_REQUEST_TIMEOUT=%r; using default %ss.", + raw, + _DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT, + ) + return _DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT + return parsed + class PrometheusLogger(CustomLogger): # Class variables or attributes @@ -1607,7 +1628,15 @@ class PrometheusLogger(CustomLogger): _user_spend = _metadata.get("user_api_key_user_spend", None) _user_max_budget = _metadata.get("user_api_key_user_max_budget", None) - results = await asyncio.gather( + # Bound the per-request budget-metric emission so that slow Redis/DB + # lookups under load cannot consume the whole LoggingWorker watchdog + # (LOGGING_WORKER_MAX_TIME_PER_COROUTINE, default 20s) and get the entire + # success-logging event cancelled. Budget gauges are also refreshed by the + # periodic cron every PROMETHEUS_BUDGET_METRICS_REFRESH_INTERVAL_MINUTES, + # so dropping one slow per-request emission only loses sub-cron real-time + # detail, not correctness. + budget_metrics_timeout = _get_budget_metrics_per_request_timeout() + gather_coro = asyncio.gather( self._set_api_key_budget_metrics_after_api_request( user_api_key=user_api_key, user_api_key_alias=user_api_key_alias, @@ -1634,6 +1663,16 @@ class PrometheusLogger(CustomLogger): ), return_exceptions=True, ) + try: + results = await asyncio.wait_for(gather_coro, timeout=budget_metrics_timeout) + except asyncio.TimeoutError: + verbose_logger.debug( + "[Non-Blocking] Prometheus: per-request budget metric emission " + "exceeded %ss under load; skipping (values are refreshed by the " + "periodic budget-metrics cron job).", + budget_metrics_timeout, + ) + return for i, r in enumerate(results): if isinstance(r, Exception): verbose_logger.debug( diff --git a/litellm/llms/anthropic/chat/guardrail_translation/handler.py b/litellm/llms/anthropic/chat/guardrail_translation/handler.py index 4506c114208..7000c20d9c4 100644 --- a/litellm/llms/anthropic/chat/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/chat/guardrail_translation/handler.py @@ -48,7 +48,10 @@ from litellm.types.utils import ( ) if TYPE_CHECKING: - from litellm.integrations.custom_guardrail import CustomGuardrail + from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + ModifyResponseException, + ) from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.types.llms.anthropic_messages.anthropic_response import ( AnthropicMessagesResponse, @@ -70,6 +73,170 @@ 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", + stream_started: bool = False, + responses_so_far: Optional[list[Any]] = None, + ) -> list[bytes]: + """ + Build an Anthropic SSE sequence delivering the guardrail block message + and terminating the stream cleanly. + + - ``stream_started`` False (buffered / pre-stream): nothing has been + sent, so emit a complete standalone message (message_start -> + content_block_* -> message_delta -> message_stop) via + FakeAnthropicMessagesStreamIterator, the same converter the + /v1/messages pre-stream block handler uses. + - ``stream_started`` True (sampling / detect-only end-of-stream): real + chunks were already sent, so *continue* the in-progress message -- + close the open content block, append the block message as a new text + block, then end the message. Emitting a second ``message_start`` here + would make Anthropic clients reject the stream. + """ + if stream_started: + return self._block_continuation_chunks(exc, responses_so_far or []) + return self._standalone_block_chunks(exc) + + def _standalone_block_chunks(self, exc: "ModifyResponseException") -> list[bytes]: + import uuid + + from litellm.llms.anthropic.experimental_pass_through.messages.fake_stream_iterator import ( + FakeAnthropicMessagesStreamIterator, + ) + from litellm.llms.base_llm.guardrail_translation.utils import ( + blocked_response_usage, + ) + from litellm.types.utils import AnthropicMessagesResponse + + block_response = AnthropicMessagesResponse( + id=f"msg_{uuid.uuid4()}", + type="message", + role="assistant", + content=[{"type": "text", "text": exc.message}], + model=exc.model, + stop_reason="end_turn", + usage=blocked_response_usage(getattr(exc, "original_response", None)), + ) + return list(FakeAnthropicMessagesStreamIterator(response=block_response)) + + def _block_continuation_chunks(self, exc: "ModifyResponseException", responses_so_far: list[Any]) -> list[bytes]: + """Continue an already-started message: close the open content block, + append the block message as a new text block, then end the message -- + without a second message_start.""" + + from litellm.llms.base_llm.guardrail_translation.utils import ( + blocked_response_usage, + ) + + def _sse(event_type: str, payload: dict) -> bytes: + return f"event: {event_type}\ndata: {json.dumps(payload)}\n\n".encode() + + output_tokens = blocked_response_usage(getattr(exc, "original_response", None))["output_tokens"] + open_index, max_index = self._content_block_state(responses_so_far) + new_index = (max_index + 1) if max_index is not None else 0 + chunks: list[bytes] = [] + if open_index is not None: + chunks.append(_sse("content_block_stop", {"type": "content_block_stop", "index": open_index})) + chunks += [ + _sse( + "content_block_start", + { + "type": "content_block_start", + "index": new_index, + "content_block": {"type": "text", "text": ""}, + }, + ), + _sse( + "content_block_delta", + { + "type": "content_block_delta", + "index": new_index, + "delta": {"type": "text_delta", "text": exc.message}, + }, + ), + _sse("content_block_stop", {"type": "content_block_stop", "index": new_index}), + _sse( + "message_delta", + { + "type": "message_delta", + "delta": {"stop_reason": "end_turn", "stop_sequence": None}, + "usage": {"output_tokens": output_tokens}, + }, + ), + _sse("message_stop", {"type": "message_stop"}), + ] + return chunks + + @staticmethod + def _content_block_state( + responses_so_far: list[Any], + ) -> tuple[Optional[int], Optional[int]]: + """From the SSE chunks already sent to the client, return (open + content-block index or None, highest content-block index seen or None). + + A single streamed item may bundle multiple SSE events (raw bytes) or be + an already-parsed event dict, so every event across every item is + considered -- matching how ``get_streaming_string_so_far`` reads the + same stream.""" + open_indices: set[int] = set() + max_index: Optional[int] = None + for item in responses_so_far: + for data in AnthropicMessagesHandler._iter_sse_events(item): + event_type = data.get("type") + index = data.get("index") + if not isinstance(index, int): + continue + if event_type == "content_block_start": + open_indices.add(index) + max_index = index if max_index is None else max(max_index, index) + elif event_type == "content_block_stop": + open_indices.discard(index) + open_index = max(open_indices) if open_indices else None + return open_index, max_index + + @staticmethod + def _iter_sse_events(item: Any) -> list[dict]: + """Yield the event-data dicts in one stream chunk. + + Handles both formats this stream can carry (see + ``get_streaming_string_so_far``): raw SSE ``bytes`` -- which may bundle + several events separated by a blank line -- and an already-parsed event + ``dict``.""" + if isinstance(item, dict): + return [item] + if not isinstance(item, (bytes, bytearray)): + return [] + events: list[dict] = [] + for block in item.decode("utf-8", errors="replace").split("\n\n"): + for line in block.split("\n"): + line = line.strip() + if not line.startswith("data:"): + continue + try: + parsed = json.loads(line[len("data:") :].strip()) + except json.JSONDecodeError: + continue + if isinstance(parsed, dict): + events.append(parsed) + return events + def _translate_to_openai(self, data: dict) -> ChatCompletionRequest: """Translate Anthropic request to OpenAI chat completion format.""" ( @@ -406,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 @@ -430,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/base_translation.py b/litellm/llms/base_llm/guardrail_translation/base_translation.py index 68db36b529e..6c41b46cfa0 100644 --- a/litellm/llms/base_llm/guardrail_translation/base_translation.py +++ b/litellm/llms/base_llm/guardrail_translation/base_translation.py @@ -2,7 +2,10 @@ from abc import ABC, abstractmethod from typing import TYPE_CHECKING, Any, Dict, List, Optional if TYPE_CHECKING: - from litellm.integrations.custom_guardrail import CustomGuardrail + from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + ModifyResponseException, + ) from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.proxy._types import UserAPIKeyAuth from litellm.types.llms.openai import AllMessageValues @@ -98,6 +101,30 @@ class BaseTranslation(ABC): """ return responses_so_far + def build_block_sse_chunks( + self, + exc: "ModifyResponseException", + stream_started: bool = False, + responses_so_far: Optional[list[Any]] = None, + ) -> Optional[list[bytes]]: + """ + Build the streaming chunks that deliver a guardrail block message and + cleanly terminate the stream in this provider's wire format. + + ``stream_started`` is True when real chunks were already sent to the + client: the result must *continue* the in-progress message (e.g. close + the open content block and append the block message) rather than start + a new one, which clients reject. ``responses_so_far`` provides the prior + chunks needed to do so. When False, nothing has been sent and a + standalone block message is emitted. + + Returns None when the format has no safe terminator; the caller then + re-raises ``exc`` so the proxy can surface a clean error instead. + Override in provider subclasses that support synthesizing a block + stream. + """ + return None + def get_structured_messages(self, data: dict) -> Optional[List["AllMessageValues"]]: """ Convert request data to OpenAI-spec structured messages. diff --git a/litellm/llms/base_llm/guardrail_translation/utils.py b/litellm/llms/base_llm/guardrail_translation/utils.py index 97ece6b5eab..8a06dd4ea52 100644 --- a/litellm/llms/base_llm/guardrail_translation/utils.py +++ b/litellm/llms/base_llm/guardrail_translation/utils.py @@ -1,10 +1,100 @@ from __future__ import annotations -from typing import Any, List +import json +from typing import Any, List, Optional +from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicUsage from litellm.types.llms.openai import AllMessageValues +def _anthropic_stream_chunk_events(item: Any) -> list[dict]: + if isinstance(item, dict): + return [item] + if isinstance(item, bytes): + chunk = item.decode("utf-8", errors="replace") + elif isinstance(item, str): + chunk = item + else: + return [] + + events: list[dict] = [] + for block in chunk.split("\n\n"): + for line in block.splitlines(): + stripped = line.strip() + if not stripped.startswith("data:"): + continue + payload = stripped[len("data:") :].strip() + if not payload or payload == "[DONE]": + continue + try: + parsed = json.loads(payload) + except json.JSONDecodeError: + continue + if isinstance(parsed, dict): + events.append(parsed) + return events + + +def _usage_from_anthropic_stream_chunks(original_response: list[Any]) -> Optional[AnthropicUsage]: + input_tokens = 0 + output_tokens = 0 + found_usage = False + + for item in original_response: + for event in _anthropic_stream_chunk_events(item): + event_type = event.get("type") + if event_type == "message_start": + message = event.get("message") or {} + usage_obj = message.get("usage") or {} + elif event_type == "message_delta": + usage_obj = event.get("usage") or {} + else: + usage_obj = {} + if not isinstance(usage_obj, dict): + continue + if usage_obj.get("input_tokens") is not None: + input_tokens = int(usage_obj.get("input_tokens") or 0) + found_usage = True + if usage_obj.get("output_tokens") is not None: + output_tokens = int(usage_obj.get("output_tokens") or 0) + found_usage = True + + if not found_usage: + return None + return AnthropicUsage(input_tokens=input_tokens, output_tokens=output_tokens) + + +def blocked_response_usage(original_response: Optional[Any]) -> AnthropicUsage: + """ + Token usage for a synthetic guardrail-blocked response. + + A post-call block replaces the LLM's response with the violation message, + but the upstream call already consumed tokens -- report that real usage + (carried on ``ModifyResponseException.original_response``) rather than + discarding it. Pre-call blocks never invoked the LLM (no original_response), + so usage is zero. + """ + usage_obj: Any = None + if isinstance(original_response, list): + stream_usage = _usage_from_anthropic_stream_chunks(original_response) + if stream_usage is not None: + return stream_usage + elif isinstance(original_response, dict): + usage_obj = original_response.get("usage") + elif original_response is not None: + usage_obj = getattr(original_response, "usage", None) + + def _tokens(key: str, fallback_key: str) -> int: + if isinstance(usage_obj, dict): + 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", "prompt_tokens"), + output_tokens=_tokens("output_tokens", "completion_tokens"), + ) + + def effective_skip_system_message_for_guardrail(guardrail_to_apply: Any) -> bool: per = getattr(guardrail_to_apply, "skip_system_message_in_guardrail", None) if per is not None: diff --git a/litellm/llms/openai/responses/guardrail_translation/handler.py b/litellm/llms/openai/responses/guardrail_translation/handler.py index 6ac33ffa44a..d1323b1a2bf 100644 --- a/litellm/llms/openai/responses/guardrail_translation/handler.py +++ b/litellm/llms/openai/responses/guardrail_translation/handler.py @@ -45,6 +45,7 @@ from litellm.types.llms.openai import ( AllMessageValues, ChatCompletionToolCallChunk, ChatCompletionToolParam, + ResponsesAPIStreamEvents, ) from litellm.types.responses.main import ( GenericResponseOutputItem, @@ -586,7 +587,14 @@ class OpenAIResponsesHandler(BaseTranslation): """ Check if the streaming has ended. """ - return all(response.choices[0].finish_reason is not None for response in responses_so_far) + if not responses_so_far: + return False + terminal_types = { + ResponsesAPIStreamEvents.RESPONSE_COMPLETED.value, + ResponsesAPIStreamEvents.RESPONSE_FAILED.value, + ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE.value, + } + return responses_so_far[-1].get("type") in terminal_types def get_streaming_string_so_far(self, responses_so_far: List[Any]) -> str: """ diff --git a/litellm/proxy/anthropic_endpoints/endpoints.py b/litellm/proxy/anthropic_endpoints/endpoints.py index 71acc1f3106..92279a9e685 100644 --- a/litellm/proxy/anthropic_endpoints/endpoints.py +++ b/litellm/proxy/anthropic_endpoints/endpoints.py @@ -12,6 +12,9 @@ from litellm.integrations.custom_guardrail import ModifyResponseException from litellm.llms.anthropic.experimental_pass_through.context_management import ( AnthropicContextManagementError, ) +from litellm.llms.base_llm.guardrail_translation.utils import ( + blocked_response_usage as _blocked_response_usage, +) from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_request_processing import ( @@ -134,6 +137,10 @@ async def anthropic_response( from litellm.types.utils import AnthropicMessagesResponse + # Report the blocked LLM response's real token usage (carried on the + # exception) instead of discarding it; zero for pre-call blocks. + _usage = _blocked_response_usage(e.original_response) + _anthropic_response = AnthropicMessagesResponse( id=f"msg_{str(uuid.uuid4())}", type="message", @@ -141,7 +148,7 @@ async def anthropic_response( content=[{"type": "text", "text": e.message}], model=e.model, stop_reason="end_turn", - usage={"input_tokens": 0, "output_tokens": 0}, + usage=_usage, ) if data.get("stream", None) is not None and data["stream"] is True: 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 79c29670d93..e0b387e92c0 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -8,7 +8,7 @@ Unified Guardrail, leveraging LiteLLM's /applyGuardrail endpoint import copy import json -from typing import Any, AsyncGenerator, List, Optional, Union +from typing import TYPE_CHECKING, Any, AsyncGenerator, List, Optional, Union from fastapi import HTTPException @@ -23,6 +23,11 @@ from litellm.proxy._types import UserAPIKeyAuth from litellm.types.guardrails import GuardrailEventHooks from litellm.types.utils import CallTypes, CallTypesLiteral +if TYPE_CHECKING: + # Imported lazily at runtime (inside the streaming hook) to avoid a + # module-level cyclic import with litellm.integrations.custom_guardrail. + from litellm.integrations.custom_guardrail import ModifyResponseException + # Call types that use NDJSON streaming (A2A); guardrail HTTPException is emitted as in-stream error A2A_CALL_TYPES = (CallTypes.asend_message, CallTypes.send_message) @@ -197,6 +202,10 @@ class UnifiedLLMGuardrails(CustomLogger): ) from litellm.types.guardrails import GuardrailEventHooks + # Local import avoids a module-level cyclic import with + # litellm.integrations.custom_guardrail. + from litellm.integrations.custom_guardrail import ModifyResponseException + guardrail_to_apply: CustomGuardrail = data.pop("guardrail_to_apply", None) if guardrail_to_apply is None: @@ -238,18 +247,51 @@ class UnifiedLLMGuardrails(CustomLogger): endpoint_translation = endpoint_guardrail_translation_mappings[CallTypes(call_type)]() - response = await endpoint_translation.process_output_response( - response=response, # type: ignore - guardrail_to_apply=guardrail_to_apply, - litellm_logging_obj=data.get("litellm_logging_obj"), - user_api_key_dict=user_api_key_dict, - request_data=data, - ) + try: + response = await endpoint_translation.process_output_response( + response=response, # type: ignore + guardrail_to_apply=guardrail_to_apply, + litellm_logging_obj=data.get("litellm_logging_obj"), + user_api_key_dict=user_api_key_dict, + request_data=data, + ) + except ModifyResponseException as e: + # The guardrail blocked the response. Attach the original LLM + # response so the endpoint handler can report its real token usage + # instead of discarding it (the block replaces the content, but the + # upstream call already consumed those tokens). + if e.original_response is None: + e.original_response = response + raise # Add guardrail to applied guardrails header add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=guardrail_to_apply.guardrail_name) return response + async def _handle_streaming_block( + self, + exc: "ModifyResponseException", + endpoint_translation: Any, + stream_started: bool, + responses_so_far: list[Any], + ) -> AsyncGenerator[Any, None]: + """ + Terminate a streamed response cleanly when a guardrail blocks it. + + Format-agnostic routing: delegates to the provider translation handler's + ``build_block_sse_chunks`` (see ``BaseTranslation.build_block_sse_chunks`` + for the ``stream_started`` / ``responses_so_far`` contract). When the + format has no safe terminator the handler returns None and we re-raise + ``exc`` so the proxy can surface a clean error. + """ + block_chunks = endpoint_translation.build_block_sse_chunks( + exc, stream_started=stream_started, responses_so_far=responses_so_far + ) + if block_chunks is None: + raise exc + for chunk in block_chunks: + yield chunk + async def async_post_call_streaming_iterator_hook( self, user_api_key_dict: UserAPIKeyAuth, @@ -271,26 +313,53 @@ class UnifiedLLMGuardrails(CustomLogger): global endpoint_guardrail_translation_mappings + # Local import avoids a module-level cyclic import with + # litellm.integrations.custom_guardrail. + from litellm.integrations.custom_guardrail import ModifyResponseException + guardrail_to_apply: CustomGuardrail = request_data.pop("guardrail_to_apply", None) - # Get streaming configuration from guardrail or optional_params - sampling_rate = 5 - end_of_stream_only = False # If True, only apply guardrail at end of stream + # Get streaming configuration. Resolution order (later wins): default + # < guardrail attribute < guardrail_config dict < this callback's + # optional_params. + def _streaming_flag(name: str, default: Any) -> Any: + value = default + if guardrail_to_apply is not None: + value = getattr(guardrail_to_apply, name, value) + config = getattr(guardrail_to_apply, "guardrail_config", {}) + if isinstance(config, dict): + value = config.get(name, value) + return self.optional_params.get(name, value) - if guardrail_to_apply is not None: - # Check direct attributes on guardrail first - sampling_rate = getattr(guardrail_to_apply, "streaming_sampling_rate", sampling_rate) - end_of_stream_only = getattr(guardrail_to_apply, "streaming_end_of_stream_only", end_of_stream_only) + sampling_rate = _streaming_flag("streaming_sampling_rate", 5) + # Only apply the guardrail at end of stream (not per chunk). + end_of_stream_only = _streaming_flag("streaming_end_of_stream_only", False) + # Withhold every chunk until end-of-stream moderation passes, then + # release the original chunks (clean) or only the block message + # (blocked) -- moderating the whole response *before* any content + # reaches the client. Only safe for allow/block guardrails: on + # release the original chunks are replayed as-is, so a + # content-rewriting guardrail (e.g. PII masking) would leak + # unredacted content. Guarded below via mask_response_content. + buffer_until_moderated = _streaming_flag("streaming_buffer_until_moderated", False) - # Also check guardrail_config dict if present - guardrail_config = getattr(guardrail_to_apply, "guardrail_config", {}) - if isinstance(guardrail_config, dict): - sampling_rate = guardrail_config.get("streaming_sampling_rate", sampling_rate) - end_of_stream_only = guardrail_config.get("streaming_end_of_stream_only", end_of_stream_only) + if ( + buffer_until_moderated + and guardrail_to_apply is not None + and getattr(guardrail_to_apply, "mask_response_content", False) + ): + verbose_proxy_logger.warning( + "UnifiedLLMGuardrails: streaming_buffer_until_moderated is disabled for %s " + "because mask_response_content=True -- buffered replay would release " + "unredacted original chunks instead of the moderated output.", + guardrail_to_apply.guardrail_name, + ) + buffer_until_moderated = False - # Also check optional_params as fallback - sampling_rate = self.optional_params.get("streaming_sampling_rate", sampling_rate) - end_of_stream_only = self.optional_params.get("streaming_end_of_stream_only", end_of_stream_only) + # Buffering can only moderate the assembled response, so it always + # defers to end-of-stream. + if buffer_until_moderated: + end_of_stream_only = True if guardrail_to_apply is None: async for item in response: @@ -315,6 +384,12 @@ 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). + chunks_yielded = False async for item in response: chunk_counter += 1 @@ -336,9 +411,22 @@ class UnifiedLLMGuardrails(CustomLogger): yield remaining_item return - # If end_of_stream_only mode, yield chunks without processing + # If end_of_stream_only mode, yield chunks without processing. + # When buffering, withhold them instead -- they are released (or + # replaced by the block message) only after end-of-stream + # moderation runs below. if end_of_stream_only: - yield item + if not buffer_until_moderated: + 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 @@ -368,6 +456,26 @@ class UnifiedLLMGuardrails(CustomLogger): user_api_key_dict=user_api_key_dict, 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 + # `data: {"error": ...}` blob (which truncates the stream). + # Chunks have already been forwarded here, so the block + # continues the in-progress message (stream_started=True). + # The current chunk was appended to responses_so_far but not + # yet yielded, so exclude it: the continuation must reflect + # only what the client has actually received. + async for block_chunk in self._handle_streaming_block( + e, + endpoint_translation, + stream_started=chunks_yielded, + responses_so_far=responses_yielded, + ): + yield block_chunk + return except HTTPException as e: # Response already started (we already yielded chunks); cannot send 400. # For A2A (NDJSON), yield an in-stream JSON-RPC error so the client sees it. @@ -394,8 +502,12 @@ class UnifiedLLMGuardrails(CustomLogger): yield error_chunk 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 @@ -408,6 +520,15 @@ class UnifiedLLMGuardrails(CustomLogger): endpoint_translation = endpoint_guardrail_translation_mappings[CallTypes(call_type)]() + # When buffering, snapshot the original chunks before moderation. + # A shallow copy suffices: end-of-stream + # process_output_streaming_response builds a separate assembled + # response (it does not mutate the individual chunks in place), and + # the chunks themselves are replayed verbatim -- so we only need to + # preserve the list, not clone every chunk (deepcopy would double + # peak memory for large responses). + buffered_items = list(responses_so_far) if buffer_until_moderated else None + try: await endpoint_translation.process_output_streaming_response( responses_so_far=responses_so_far, @@ -416,6 +537,28 @@ class UnifiedLLMGuardrails(CustomLogger): user_api_key_dict=user_api_key_dict, request_data=request_data, ) + # Moderation passed: release the withheld original chunks. + 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. + # The withheld original chunks are never released. + async for block_chunk in self._handle_streaming_block( + e, + endpoint_translation, + stream_started=bool(responses_yielded), + responses_so_far=responses_yielded, + ): + yield block_chunk + return except HTTPException as e: if call_type is not None and CallTypes(call_type) in A2A_CALL_TYPES: request_id = _get_a2a_request_id(responses_so_far, request_data) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index a29acf3b06f..bb925799bb8 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -8361,6 +8361,22 @@ async def model_info( ) +def _blocked_response_usage(original_response: Optional[Any]) -> "litellm.Usage": + """ + Token usage for a synthetic guardrail-blocked response. + + A post-call block replaces the LLM's response with the violation message, + but the upstream call already consumed tokens -- report that real usage + (carried on ``ModifyResponseException.original_response``) rather than + discarding it. Pre-call blocks never invoked the LLM (no original_response), + so usage is zero. + """ + usage = getattr(original_response, "usage", None) if original_response is not None else None + if isinstance(usage, litellm.Usage): + return usage + return litellm.Usage(prompt_tokens=0, completion_tokens=0, total_tokens=0) + + @router.post( "/v1/chat/completions", dependencies=[Depends(user_api_key_auth)], @@ -8467,6 +8483,9 @@ async def chat_completion( _chat_response.model = e.model # type: ignore _chat_response.choices[0].message.content = e.message # type: ignore _chat_response.choices[0].finish_reason = "content_filter" # type: ignore + # Report the blocked LLM response's real usage (set before the stream + # branch so both paths carry it); zero for pre-call blocks. + _chat_response.usage = _blocked_response_usage(e.original_response) # type: ignore if data.get("stream", None) is not None and data["stream"] is True: _iterator = litellm.utils.ModelResponseIterator(model_response=_chat_response, convert_to_delta=True) @@ -8488,8 +8507,6 @@ async def chat_completion( media_type="text/event-stream", status_code=200, # Return 200 for passthrough mode ) - _usage = litellm.Usage(prompt_tokens=0, completion_tokens=0, total_tokens=0) - _chat_response.usage = _usage # type: ignore return _chat_response except RejectedRequestError as e: _data = e.request_data @@ -8618,11 +8635,7 @@ async def completion( # Set text attribute dynamically for text completion format setattr(_text_response.choices[0], "text", e.message) _text_response.model = e.model # type: ignore[assignment] - _usage = litellm.Usage( - prompt_tokens=0, - completion_tokens=0, - total_tokens=0, - ) + _usage = _blocked_response_usage(e.original_response) # Set usage attribute dynamically (ModelResponse accepts usage in __init__ but it's not in type definition) setattr(_text_response, "usage", _usage) _iterator = litellm.utils.ModelResponseIterator(model_response=_text_response, convert_to_delta=True) @@ -8647,11 +8660,7 @@ async def completion( _response = litellm.TextCompletionResponse() _response.choices[0].text = e.message _response.model = e.model # type: ignore - _usage = litellm.Usage( - prompt_tokens=0, - completion_tokens=0, - total_tokens=0, - ) + _usage = _blocked_response_usage(e.original_response) _response.usage = _usage # type: ignore return _response except RejectedRequestError as e: diff --git a/tests/test_litellm/integrations/test_prometheus_budget_metrics_timeout.py b/tests/test_litellm/integrations/test_prometheus_budget_metrics_timeout.py new file mode 100644 index 00000000000..a4d245e9dc0 --- /dev/null +++ b/tests/test_litellm/integrations/test_prometheus_budget_metrics_timeout.py @@ -0,0 +1,168 @@ +""" +Unit tests for the per-request budget-metric emission timeout in +PrometheusLogger._increment_remaining_budget_metrics. + +A slow Redis/DB lookup in one of the budget branches must not let the gather run +unbounded; it is wrapped in asyncio.wait_for so the success-logging coroutine +cannot exceed the LoggingWorker watchdog and get the whole event cancelled. +""" + +import asyncio +from unittest.mock import AsyncMock, patch + +import pytest +from prometheus_client import REGISTRY + +from litellm.integrations.prometheus import ( + PrometheusLogger, + _DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT, + _get_budget_metrics_per_request_timeout, +) + +TIMEOUT_ENV = "PROMETHEUS_BUDGET_METRICS_PER_REQUEST_TIMEOUT" + + +@pytest.fixture(autouse=True) +def cleanup_prometheus_registry(): + collectors = list(REGISTRY._collector_to_names.keys()) + for collector in collectors: + try: + REGISTRY.unregister(collector) + except Exception: + pass + yield + collectors = list(REGISTRY._collector_to_names.keys()) + for collector in collectors: + try: + REGISTRY.unregister(collector) + except Exception: + pass + + +@pytest.fixture +def prometheus_logger(): + return PrometheusLogger() + + +def _call_increment(logger: PrometheusLogger): + return logger._increment_remaining_budget_metrics( + user_api_team="team-1", + user_api_team_alias="team-alias", + user_api_key="key-1", + user_api_key_alias="key-alias", + litellm_params={"metadata": {}}, + response_cost=0.01, + user_id="user-1", + user_api_key_org_id="org-1", + ) + + +def _skip_logged(debug_mock) -> bool: + return any("skipping" in str(call.args[0]) for call in debug_mock.call_args_list if call.args) + + +@pytest.mark.asyncio +async def test_budget_metric_emission_skips_on_timeout(prometheus_logger, monkeypatch): + """A branch slower than the timeout is skipped without propagating, and the + skip is logged instead of cancelling the success-logging event.""" + monkeypatch.setenv(TIMEOUT_ENV, "0.05") + + async def _slow_branch(**kwargs): + await asyncio.sleep(30) + + prometheus_logger._set_api_key_budget_metrics_after_api_request = _slow_branch + prometheus_logger._set_team_budget_metrics_after_api_request = AsyncMock() + prometheus_logger._set_user_budget_metrics_after_api_request = AsyncMock() + prometheus_logger._set_org_budget_metrics_after_api_request = AsyncMock() + + with patch("litellm.integrations.prometheus.verbose_logger") as mock_logger: + await _call_increment(prometheus_logger) + + assert _skip_logged(mock_logger.debug) + + +@pytest.mark.asyncio +async def test_budget_metric_emission_completes_within_timeout(prometheus_logger, monkeypatch): + """With a generous timeout every branch is awaited and no skip is logged.""" + monkeypatch.setenv(TIMEOUT_ENV, "5.0") + + prometheus_logger._set_api_key_budget_metrics_after_api_request = AsyncMock() + prometheus_logger._set_team_budget_metrics_after_api_request = AsyncMock() + prometheus_logger._set_user_budget_metrics_after_api_request = AsyncMock() + prometheus_logger._set_org_budget_metrics_after_api_request = AsyncMock() + + with patch("litellm.integrations.prometheus.verbose_logger") as mock_logger: + await _call_increment(prometheus_logger) + + assert prometheus_logger._set_api_key_budget_metrics_after_api_request.await_count == 1 + assert prometheus_logger._set_team_budget_metrics_after_api_request.await_count == 1 + assert prometheus_logger._set_user_budget_metrics_after_api_request.await_count == 1 + assert prometheus_logger._set_org_budget_metrics_after_api_request.await_count == 1 + assert not _skip_logged(mock_logger.debug) + + +@pytest.mark.asyncio +async def test_invalid_timeout_env_falls_back_to_default(prometheus_logger, monkeypatch): + """A malformed timeout env value must not raise (which would recreate the + failure mode); it falls back to the default and every branch still runs.""" + monkeypatch.setenv(TIMEOUT_ENV, "not-a-number") + + prometheus_logger._set_api_key_budget_metrics_after_api_request = AsyncMock() + prometheus_logger._set_team_budget_metrics_after_api_request = AsyncMock() + prometheus_logger._set_user_budget_metrics_after_api_request = AsyncMock() + prometheus_logger._set_org_budget_metrics_after_api_request = AsyncMock() + + await _call_increment(prometheus_logger) + + assert prometheus_logger._set_api_key_budget_metrics_after_api_request.await_count == 1 + assert prometheus_logger._set_org_budget_metrics_after_api_request.await_count == 1 + + +@pytest.mark.parametrize("value", ["not-a-number", "0", "-1", "nan", "inf", "-inf"]) +def test_unusable_timeout_env_falls_back_to_default(value, monkeypatch): + """Values that parse but disable or unbound the timeout (0, negative, nan, + inf) must fall back to the default instead of being used; otherwise they + either skip every emission or recreate the unbounded-wait failure mode.""" + monkeypatch.setenv(TIMEOUT_ENV, value) + + assert _get_budget_metrics_per_request_timeout() == _DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT + + +@pytest.mark.parametrize("value,expected", [("0.05", 0.05), ("5.0", 5.0), ("30", 30.0)]) +def test_valid_timeout_env_is_used(value, expected, monkeypatch): + """A finite positive value is parsed and returned unchanged.""" + monkeypatch.setenv(TIMEOUT_ENV, value) + + assert _get_budget_metrics_per_request_timeout() == expected + + +def test_missing_timeout_env_uses_default(monkeypatch): + """With the env unset the default is returned.""" + monkeypatch.delenv(TIMEOUT_ENV, raising=False) + + assert _get_budget_metrics_per_request_timeout() == _DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT + + +@pytest.mark.asyncio +async def test_outer_cancellation_still_propagates(prometheus_logger, monkeypatch): + """Only asyncio.TimeoutError is swallowed; an outer cancellation (cooperative + shutdown / watchdog) injected while awaiting must still propagate.""" + monkeypatch.setenv(TIMEOUT_ENV, "30") + + started = asyncio.Event() + + async def _slow_branch(**kwargs): + started.set() + await asyncio.sleep(30) + + prometheus_logger._set_api_key_budget_metrics_after_api_request = _slow_branch + prometheus_logger._set_team_budget_metrics_after_api_request = AsyncMock() + prometheus_logger._set_user_budget_metrics_after_api_request = AsyncMock() + prometheus_logger._set_org_budget_metrics_after_api_request = AsyncMock() + + task = asyncio.create_task(_call_increment(prometheus_logger)) + await started.wait() + task.cancel() + + with pytest.raises(asyncio.CancelledError): + await task diff --git a/tests/test_litellm/proxy/anthropic_endpoints/test_endpoints.py b/tests/test_litellm/proxy/anthropic_endpoints/test_endpoints.py index a4da4587b7f..69d90a8b59b 100644 --- a/tests/test_litellm/proxy/anthropic_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/anthropic_endpoints/test_endpoints.py @@ -58,15 +58,71 @@ class TestAnthropicEndpoints(unittest.TestCase): self.assertEqual(result, expected_result) # Assert safe_dumps was called for dictionary objects - mock_safe_dumps.assert_any_call( - {"type": "message_start", "message": {"id": "msg_123"}} + mock_safe_dumps.assert_any_call({"type": "message_start", "message": {"id": "msg_123"}}) + mock_safe_dumps.assert_any_call({"type": "content_block_delta", "delta": {"text": "more data"}}) + assert mock_safe_dumps.call_count == 2 # Called twice, once for each dict object + + +class TestBlockedResponseUsage: + """Blocked responses report the blocked LLM response's real usage.""" + + def test_uses_original_response_usage(self): + from litellm.proxy.anthropic_endpoints.endpoints import _blocked_response_usage + + # original_response is the AnthropicMessagesResponse the LLM produced + # before the guardrail blocked it; its usage is real. + original = {"usage": {"input_tokens": 31, "output_tokens": 9}} + assert _blocked_response_usage(original) == { + "input_tokens": 31, + "output_tokens": 9, + } + + def test_zero_usage_when_no_original_response(self): + from litellm.proxy.anthropic_endpoints.endpoints import _blocked_response_usage + + # Pre-call blocks never invoked the LLM -> nothing consumed. + assert _blocked_response_usage(None) == { + "input_tokens": 0, + "output_tokens": 0, + } + + @pytest.mark.asyncio + async def test_blocked_endpoint_response_carries_original_usage(self): + """The /v1/messages block handler reports the blocked response's real + usage, carried on ModifyResponseException.original_response.""" + from unittest.mock import AsyncMock, MagicMock + + import litellm.proxy.anthropic_endpoints.endpoints as ep + import litellm.proxy.proxy_server as proxy_server + from litellm.integrations.custom_guardrail import ModifyResponseException + + exc = ModifyResponseException( + message="blocked by guardrail", + model="claude-3-5-sonnet-20240620", + request_data={"messages": [{"role": "user", "content": "hi"}]}, + guardrail_name="rubrik", + original_response={"usage": {"input_tokens": 12, "output_tokens": 5}}, ) - mock_safe_dumps.assert_any_call( - {"type": "content_block_delta", "delta": {"text": "more data"}} - ) - assert ( - mock_safe_dumps.call_count == 2 - ) # Called twice, once for each dict object + + with ( + patch.object(ep, "_read_request_body", new=AsyncMock(return_value={})), + patch.object( + ep.ProxyBaseLLMRequestProcessing, + "base_process_llm_request", + new=AsyncMock(side_effect=exc), + ), + patch.object(proxy_server, "proxy_logging_obj") as mock_logging, + ): + mock_logging.post_call_failure_hook = AsyncMock() + response = await ep.anthropic_response( + fastapi_response=MagicMock(), + request=MagicMock(), + user_api_key_dict=MagicMock(), + ) + + assert response["content"][0]["text"] == "blocked by guardrail" + assert response["usage"] == {"input_tokens": 12, "output_tokens": 5} + mock_logging.post_call_failure_hook.assert_awaited_once() class TestEventLoggingBatchEndpoint: @@ -159,9 +215,7 @@ class TestStripTotalTokens(unittest.TestCase): # SimpleNamespace mimics the .usage attribute access pattern; the # helper's contract: if .usage is dict-shaped, strip total_tokens. - response = SimpleNamespace( - usage={"input_tokens": 100, "output_tokens": 50, "total_tokens": 150} - ) + response = SimpleNamespace(usage={"input_tokens": 100, "output_tokens": 50, "total_tokens": 150}) _strip_total_tokens_from_anthropic_response(response) assert "total_tokens" not in response.usage assert response.usage == {"input_tokens": 100, "output_tokens": 50} 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 new file mode 100644 index 00000000000..865b4164e5e --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_anthropic_streaming_block.py @@ -0,0 +1,362 @@ +""" +Regression tests for blocking an Anthropic streaming response from the +unified guardrail post-call streaming iterator hook. + +When a guardrail's ``apply_guardrail`` raises ``ModifyResponseException`` while +(or at the end of) an Anthropic ``/v1/messages`` stream is being relayed, the +hook must emit a well-formed Anthropic SSE termination sequence carrying the +block message - NOT a bare ``data: {"error": ...}`` blob that truncates the +stream and causes the Anthropic SDK parser to discard the response. +""" + +import json +from typing import Any, List, Literal, Optional + +import pytest + +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + ModifyResponseException, +) +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import ( + UnifiedLLMGuardrails, +) +from litellm.types.utils import GenericGuardrailAPIInputs + +BLOCK_MESSAGE = "Blocked by policy: this response was withheld." + + +class _BlockingGuardrail(CustomGuardrail): + """Mock guardrail that always blocks by raising ModifyResponseException.""" + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional[Any] = None, + ) -> GenericGuardrailAPIInputs: + raise ModifyResponseException( + message=BLOCK_MESSAGE, + model="claude-3-5-sonnet", + request_data=request_data, + guardrail_name=self.guardrail_name, + ) + + +def _sse_event(event_type: str, data: dict) -> bytes: + return f"event: {event_type}\ndata: {json.dumps(data)}\n\n".encode() + + +async def _anthropic_stream(end: bool): + """Yield Anthropic SSE byte chunks. If end=True, include a terminating + message_delta (stop_reason set) so the hook's end-of-stream path runs.""" + yield _sse_event( + "message_start", + { + "type": "message_start", + "message": { + "id": "msg_orig", + "type": "message", + "role": "assistant", + "model": "claude-3-5-sonnet", + "content": [], + "stop_reason": None, + "usage": {"input_tokens": 1, "output_tokens": 0}, + }, + }, + ) + yield _sse_event( + "content_block_start", + { + "type": "content_block_start", + "index": 0, + "content_block": {"type": "text", "text": ""}, + }, + ) + for text in ["This ", "is ", "the ", "original ", "answer."]: + yield _sse_event( + "content_block_delta", + { + "type": "content_block_delta", + "index": 0, + "delta": {"type": "text_delta", "text": text}, + }, + ) + if end: + yield _sse_event("content_block_stop", {"type": "content_block_stop", "index": 0}) + yield _sse_event( + "message_delta", + { + "type": "message_delta", + "delta": {"stop_reason": "end_turn", "stop_sequence": None}, + "usage": {"output_tokens": 5}, + }, + ) + yield _sse_event("message_stop", {"type": "message_stop"}) + + +def _decode(chunks: List[Any]) -> str: + parts = [] + for chunk in chunks: + parts.append(chunk.decode() if isinstance(chunk, bytes) else str(chunk)) + return "".join(parts) + + +def _parse_sse_event_types(raw: str) -> List[str]: + event_types = [] + for block in raw.split("\n\n"): + for line in block.strip().split("\n"): + if line.startswith("data:"): + payload = line[len("data:") :].strip() + try: + event_types.append(json.loads(payload).get("type")) + except json.JSONDecodeError: + pass + return event_types + + +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") + request_data = { + "messages": [{"role": "user", "content": "hi"}], + "guardrail_to_apply": guardrail, + "metadata": {"guardrails": ["test-blocking-guardrail"]}, + } + + collected: List[Any] = [] + async for chunk in unified_guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=user_api_key_dict, + response=_anthropic_stream(end=end), + request_data=request_data, + ): + collected.append(chunk) + return _decode(collected) + + +def _assert_clean_block_termination(raw: str) -> None: + # No bare error blob that would truncate the stream. + assert '"error"' not in raw, f"unexpected error blob in stream: {raw!r}" + # The block message is delivered as assistant text. + assert BLOCK_MESSAGE in raw, f"block message missing from stream: {raw!r}" + # A complete, parseable Anthropic SSE termination sequence is present. + event_types = _parse_sse_event_types(raw) + assert "message_start" in event_types + assert "content_block_delta" in event_types + # Exactly one message_start: a block must never inject a second + # message envelope into an already-started stream (clients reject it). + assert event_types.count("message_start") == 1, f"expected a single message_start, got: {event_types}" + assert event_types[-1] == "message_stop", f"stream did not end cleanly: {event_types}" + # message_delta carries a stop_reason. + 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.""" + raw = await _run_hook(end=False) + _assert_clean_block_termination(raw) + + +@pytest.mark.asyncio +async def test_end_of_stream_block_emits_clean_anthropic_sse(): + """End-of-stream block: same clean SSE termination guarantees.""" + raw = await _run_hook(end=True) + _assert_clean_block_termination(raw) + + +@pytest.mark.asyncio +async def test_mid_stream_block_after_prior_chunks_continues_message(): + """Regression: when real chunks were already forwarded (sampling_rate>1), + the block must continue the in-progress message, not start a second one.""" + raw = await _run_hook(end=False, sampling_rate=5) + # Some original content was forwarded before the block... + assert "message_start" in raw + # ...and the block continues that same message (single message_start) with + # the block message appended, ending cleanly. + _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 + + +def test_blocked_stream_reports_usage_from_original_chunks(): + from litellm.integrations.custom_guardrail import ModifyResponseException + from litellm.llms.anthropic.chat.guardrail_translation.handler import ( + AnthropicMessagesHandler, + ) + from litellm.llms.base_llm.guardrail_translation.utils import ( + blocked_response_usage, + ) + + original_chunks: List[Any] = [ + _sse_event( + "message_start", + { + "type": "message_start", + "message": { + "id": "msg_orig", + "type": "message", + "role": "assistant", + "model": "claude-3-5-sonnet", + "content": [], + "stop_reason": None, + "usage": {"input_tokens": 12, "output_tokens": 0}, + }, + }, + ), + {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 5}}, + ] + seen_chunks = [ + _sse_event( + "content_block_start", + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + ) + ] + exc = ModifyResponseException( + message=BLOCK_MESSAGE, + model="claude-3-5-sonnet", + request_data={}, + guardrail_name="g", + original_response=original_chunks, + ) + + usage = blocked_response_usage(original_chunks) + raw = b"".join( + AnthropicMessagesHandler().build_block_sse_chunks(exc, stream_started=True, responses_so_far=seen_chunks) + ).decode() + message_delta_usages = [ + payload.get("usage", {}).get("output_tokens") + for payload in _parse_sse_payloads(raw) + if payload.get("type") == "message_delta" + ] + + assert usage == {"input_tokens": 12, "output_tokens": 5} + 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), + so a mid-stream block closes/opens the right indices.""" + + def _handler(self): + from litellm.llms.anthropic.chat.guardrail_translation.handler import ( + AnthropicMessagesHandler, + ) + + return AnthropicMessagesHandler() + + def test_multi_event_bytes_chunk_is_fully_parsed(self): + # One item bundles start(0) + delta + stop(0): the block is already + # closed, so open_index is None (not 0) and max_index is 0. + bundled = ( + _sse_event( + "content_block_start", + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + ) + + _sse_event( + "content_block_delta", + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "hi"}}, + ) + + _sse_event("content_block_stop", {"type": "content_block_stop", "index": 0}) + ) + open_index, max_index = self._handler()._content_block_state([bundled]) + assert open_index is None + assert max_index == 0 + + def test_open_block_across_separate_chunks(self): + chunks = [ + _sse_event( + "content_block_start", + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + ), + _sse_event("content_block_stop", {"type": "content_block_stop", "index": 0}), + _sse_event( + "content_block_start", + {"type": "content_block_start", "index": 1, "content_block": {"type": "text", "text": ""}}, + ), + ] + open_index, max_index = self._handler()._content_block_state(chunks) + assert open_index == 1 + assert max_index == 1 + + def test_dict_format_chunks_are_parsed(self): + # The backwards-compat parsed-dict format must be understood too. + chunks = [ + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "x"}}, + ] + open_index, max_index = self._handler()._content_block_state(chunks) + assert open_index == 0 + assert max_index == 0 + + def test_continuation_closes_open_block_and_appends_after_it(self): + from litellm.integrations.custom_guardrail import ModifyResponseException + + handler = self._handler() + # Client has seen an open text block at index 0. + seen = [ + _sse_event( + "content_block_start", + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + ), + _sse_event( + "content_block_delta", + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "partial"}}, + ), + ] + exc = ModifyResponseException( + message=BLOCK_MESSAGE, model="claude-3-5-sonnet", request_data={}, guardrail_name="g" + ) + raw = b"".join(handler.build_block_sse_chunks(exc, stream_started=True, responses_so_far=seen)).decode() + events = _parse_sse_event_types(raw) + # No new message envelope, closes block 0, appends block text at index 1. + assert "message_start" not in events + assert events == [ + "content_block_stop", + "content_block_start", + "content_block_delta", + "content_block_stop", + "message_delta", + "message_stop", + ] + assert BLOCK_MESSAGE in raw + assert '"index": 1' in raw diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_streaming_buffer_until_moderated.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_streaming_buffer_until_moderated.py new file mode 100644 index 00000000000..2b163ee5233 --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_streaming_buffer_until_moderated.py @@ -0,0 +1,173 @@ +""" +Tests for ``streaming_buffer_until_moderated`` on the unified guardrail +post-call streaming iterator hook. + +With this flag set, the hook must withhold every upstream chunk until +end-of-stream moderation has run. The decisive guarantee versus the +detect-only ``streaming_end_of_stream_only`` behavior: when the guardrail +blocks, the original (objectionable) content is NEVER yielded to the client -- +only the block message is. On a clean response, all original chunks are +released unchanged after moderation passes. +""" + +import json +from typing import Any, List, Literal, Optional + +import pytest + +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + ModifyResponseException, +) +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import ( + UnifiedLLMGuardrails, +) +from litellm.types.utils import GenericGuardrailAPIInputs + +BLOCK_MESSAGE = "Blocked by policy: this response was withheld." +ORIGINAL_MARKER = "ORIGINAL-SECRET-ANSWER" + + +class _BlockingGuardrail(CustomGuardrail): + """Always blocks at moderation time.""" + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional[Any] = None, + ) -> GenericGuardrailAPIInputs: + raise ModifyResponseException( + message=BLOCK_MESSAGE, + model="claude-3-5-sonnet", + request_data=request_data, + guardrail_name=self.guardrail_name, + ) + + +class _PassingGuardrail(CustomGuardrail): + """Never blocks; returns inputs unchanged.""" + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional[Any] = None, + ) -> GenericGuardrailAPIInputs: + return inputs + + +def _sse_event(event_type: str, data: dict) -> bytes: + return f"event: {event_type}\ndata: {json.dumps(data)}\n\n".encode() + + +async def _anthropic_stream(): + """A complete Anthropic /v1/messages SSE stream whose assistant text + contains ORIGINAL_MARKER so leakage is unambiguous to assert.""" + yield _sse_event( + "message_start", + { + "type": "message_start", + "message": { + "id": "msg_orig", + "type": "message", + "role": "assistant", + "model": "claude-3-5-sonnet", + "content": [], + "stop_reason": None, + "usage": {"input_tokens": 1, "output_tokens": 0}, + }, + }, + ) + yield _sse_event( + "content_block_start", + { + "type": "content_block_start", + "index": 0, + "content_block": {"type": "text", "text": ""}, + }, + ) + for text in ["Here is ", "the ", ORIGINAL_MARKER, " for you."]: + yield _sse_event( + "content_block_delta", + { + "type": "content_block_delta", + "index": 0, + "delta": {"type": "text_delta", "text": text}, + }, + ) + yield _sse_event("content_block_stop", {"type": "content_block_stop", "index": 0}) + yield _sse_event( + "message_delta", + { + "type": "message_delta", + "delta": {"stop_reason": "end_turn", "stop_sequence": None}, + "usage": {"output_tokens": 5}, + }, + ) + yield _sse_event("message_stop", {"type": "message_stop"}) + + +def _decode(chunks: List[Any]) -> str: + return "".join(c.decode() if isinstance(c, bytes) else str(c) for c in chunks) + + +async def _run(guardrail: CustomGuardrail) -> str: + # Rubrik's real config: end-of-stream-only moderation. Without buffering + # this releases every chunk before moderation runs (content leaks on + # block); the buffer flag must change that to moderate-then-release. + guardrail.streaming_end_of_stream_only = True + guardrail.streaming_buffer_until_moderated = True + unified = UnifiedLLMGuardrails() + user_api_key_dict = UserAPIKeyAuth(api_key="test", request_route="/v1/messages") + request_data = { + "messages": [{"role": "user", "content": "hi"}], + "guardrail_to_apply": guardrail, + "metadata": {"guardrails": [guardrail.guardrail_name]}, + } + collected: List[Any] = [] + async for chunk in unified.async_post_call_streaming_iterator_hook( + user_api_key_dict=user_api_key_dict, + response=_anthropic_stream(), + request_data=request_data, + ): + collected.append(chunk) + return _decode(collected) + + +@pytest.mark.asyncio +async def test_buffered_block_withholds_original_content(): + raw = await _run(_BlockingGuardrail(guardrail_name="blk", event_hook="post_call")) + # The original content must never reach the client... + assert ORIGINAL_MARKER not in raw, f"original content leaked: {raw!r}" + # ...only the block message, in a clean terminating stream. + assert BLOCK_MESSAGE in raw + assert '"error"' not in raw + + +@pytest.mark.asyncio +async def test_buffered_clean_releases_all_content(): + raw = await _run(_PassingGuardrail(guardrail_name="pass", event_hook="post_call")) + # A clean response is released in full after moderation passes. + assert ORIGINAL_MARKER in raw + assert ( + raw.rstrip().endswith('event: message_stop\ndata: {"type": "message_stop"}'.rstrip()) or "message_stop" in raw + ) + assert BLOCK_MESSAGE not in raw + + +@pytest.mark.asyncio +async def test_buffered_mode_disabled_for_content_rewriting_guardrail(): + """Buffered replay yields the withheld *original* chunks verbatim, which + is unsafe for a guardrail that rewrites response text (e.g. PII masking): + the client would get the unredacted original instead of the moderated + output. mask_response_content=True must force buffering off so the + request falls back to the (correctly moderated) non-buffered path.""" + guardrail = _PassingGuardrail(guardrail_name="masker", event_hook="post_call", mask_response_content=True) + raw = await _run(guardrail) + assert guardrail.streaming_buffer_until_moderated is True # request asked for buffering + assert ORIGINAL_MARKER in raw + assert BLOCK_MESSAGE not in raw diff --git a/tests/test_litellm/proxy/test_blocked_response_usage.py b/tests/test_litellm/proxy/test_blocked_response_usage.py new file mode 100644 index 00000000000..d486431ca3e --- /dev/null +++ b/tests/test_litellm/proxy/test_blocked_response_usage.py @@ -0,0 +1,84 @@ +""" +Token usage on synthetic guardrail-blocked responses for the OpenAI-format +proxy endpoints (/v1/chat/completions and /v1/completions). + +A post-call block replaces the LLM response with the violation message, but the +upstream call already consumed tokens. `_blocked_response_usage` reports that +real usage (carried on `ModifyResponseException.original_response`) rather than +zero; a pre-call block never invoked the LLM, so usage is zero. +""" + +import pytest + +import litellm +from litellm.proxy.proxy_server import _blocked_response_usage + + +def test_uses_original_response_usage(): + resp = litellm.ModelResponse() + resp.usage = litellm.Usage(prompt_tokens=42, completion_tokens=7, total_tokens=49) + + usage = _blocked_response_usage(resp) + + assert usage.prompt_tokens == 42 + assert usage.completion_tokens == 7 + assert usage.total_tokens == 49 + + +def test_zero_usage_when_no_original_response(): + usage = _blocked_response_usage(None) + + assert usage.prompt_tokens == 0 + assert usage.completion_tokens == 0 + assert usage.total_tokens == 0 + + +@pytest.mark.asyncio +async def test_success_hook_attaches_original_response_on_block(): + """The unified guardrail's post-call success hook must attach the blocked + LLM response to ModifyResponseException so its real usage isn't discarded.""" + from unittest.mock import AsyncMock, MagicMock, patch + + import litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail as ug + from litellm.integrations.custom_guardrail import ModifyResponseException + from litellm.proxy._types import UserAPIKeyAuth + from litellm.types.utils import CallTypes + + response = litellm.ModelResponse() + response.usage = litellm.Usage(prompt_tokens=15, completion_tokens=3, total_tokens=18) + + guardrail = MagicMock() + guardrail.should_run_guardrail.return_value = True + guardrail.guardrail_name = "rubrik" + + # The translation layer raises a block without pre-setting original_response. + translation = MagicMock() + translation.process_output_response = AsyncMock( + side_effect=ModifyResponseException( + message="blocked", + model="gpt-4o", + request_data={}, + guardrail_name="rubrik", + ) + ) + + unified = ug.UnifiedLLMGuardrails() + user_api_key_dict = UserAPIKeyAuth(api_key="test", request_route="/chat/completions") + data = {"guardrail_to_apply": guardrail, "model": "gpt-4o"} + + # Inject our translation for the inferred call type (the module global is + # cached across tests, so patch it directly rather than the loader). + with patch.object( + ug, + "endpoint_guardrail_translation_mappings", + { + CallTypes.acompletion: lambda: translation, + CallTypes.completion: lambda: translation, + }, + ): + with pytest.raises(ModifyResponseException) as excinfo: + await unified.async_post_call_success_hook( + data=data, user_api_key_dict=user_api_key_dict, response=response + ) + + assert excinfo.value.original_response is response