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) <noreply@anthropic.com>

* 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) <noreply@anthropic.com>

* 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 <fizar@me.com>
Co-authored-by: Joseph Barker <156112794+seph-barker@users.noreply.github.com>
Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
Co-authored-by: Cursor Agent <cursoragent@cursor.com>
This commit is contained in:
Sameer Kankute 2026-07-03 09:27:31 +05:30 • committed by GitHub
parent 30ddef78d9
commit 321345d4c8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
14 changed files with 1412 additions and 63 deletions

View file

@ -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)

View file

@ -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(

View file

@ -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(

View file

@ -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.

View file

@ -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:

View file

@ -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:
"""

View file

@ -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:

View file

@ -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)

View file

@ -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:

View file

@ -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

View file

@ -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}

View file

@ -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

View file

@ -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

View file

@ -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