fix(guardrails): run the end-of-stream post_call scan when the client disconnects mid-stream (#43839)

* fix(guardrails): run end-of-stream post_call scan when the client disconnects mid-stream

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): close the guardrail stream chain in async_data_generator on client disconnect

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): leave the raw upstream response to the shielded finalizer on client disconnect

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): keep disconnect cleanup going when a streaming callback cleanup raises

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): assert the refund through a recorder instead of the mock

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(guardrails): inspect tool calls released before a disconnect under incremental_diff and record a failed scan marker

The incremental_diff transform stream now scans tool calls it already released when the client disconnects, and a disconnect scan whose translation raises after the guardrail recorded success also records guardrail_failed_to_respond

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(guardrails): pin that a text-only disconnect scan is not handed a tool_calls finish

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(guardrails): cover disconnect scans on every streaming endpoint and client, plus outage, worker-kill and cache-hit cells

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(guardrails): prove the cache-hit twin is served from the cache

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(guardrails): type the disconnect-close streams so basedpyright stops reporting unknown arguments

* fix(guardrails): give the guardrail metadata cast a reason so the type discipline gate accepts it

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(guardrails): scan released Messages and Responses tool calls on disconnect

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* style(guardrails): format the disconnect scan unit tests

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* chore(guardrails): drop mutable-ok markers that no longer suppress a rule

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(guardrails): scan released Responses output after a finished item and end only in-flight Chat choices

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(guardrails): type the disconnect scan test helpers

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(guardrails): type request_data in the disconnect scan helpers

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(guardrails): type request_data in the disconnect scan test doubles

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(guardrails): pin that chat streams with no tool call in flight end as released

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(guardrails): scan only released chunks on disconnect and skip it once a block owns the verdict

The disconnect scan now uses the chunks actually yielded to the client, copies them before scanning,
skips when a mid-stream block or HTTP error already settled the verdict, and the iterator wrapper
only closes hooks that are async generators so plain async iterator hooks keep working

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(guardrails): pin that a delivered guardrail error or final chunk settles the disconnect verdict

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(guardrails): close any hook iterator that exposes aclose when the stream ends early

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): accept a synchronous aclose on custom streaming hook iterators

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): swallow callback aclose errors at end of stream

A custom callback whose async_post_call_streaming_iterator_hook returns a
non-generator async iterator with a raising aclose() failed the finished
stream: content plus usage reached the client and then the stream surfaced
an error SSE with no [DONE], or aborted a post_call pipeline's buffering
loop into a 500 with an empty body. Wrap the aclose invocation in
_wrap_streaming_iterator_with_enrichment in try/except and log a warning
naming the callback and the cleanup error, matching close_guarded_stream
and _close_guarded_layers. Iteration-time hook exceptions still propagate.

* fix(proxy): log only the error type when a callback aclose raises

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: yucheng <yucheng@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-05 01:44:43 -07:00 • committed by GitHub
parent 1238cfe90f
commit 7a7d27c550
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
19 changed files with 2527 additions and 229 deletions

View file

@ -224,6 +224,11 @@ def _write_back_message_text(message: _WritableMessage, target: MessageTextTarge
_TOOL_USE_INPUT_ADAPTER: Final = TypeAdapter(dict[str, object])
_RELEASED_TOOL_USE_STOP: Final = (
b"event: message_delta\n"
b'data: {"type": "message_delta", "delta": {"stop_reason": "tool_use", "stop_sequence": null}, '
b'"usage": {"output_tokens": 0}}\n\n'
)
def _rewritten_tool_use_input(arguments: str) -> Mapping[str, object] | None:
@ -1573,6 +1578,12 @@ class AnthropicMessagesHandler(BaseTranslation):
tool_calls_in_flight=bool(tool_use_fingerprints) and not stream_ended,
)
def released_stream_as_ended(self, responses_so_far: Sequence[object]) -> tuple[object, ...]:
released_key: Final = self.get_streaming_scan_key(responses_so_far)
if released_key is None or not released_key.tool_calls_in_flight:
return tuple(responses_so_far)
return (*responses_so_far, _RELEASED_TOOL_USE_STOP)
@classmethod
def _streamed_tool_use_fingerprints(cls, responses_so_far: Sequence[object]) -> tuple[str, ...]:
return tuple(

View file

@ -203,6 +203,11 @@ class BaseTranslation(ABC):
def get_streaming_scan_key(self, responses_so_far: Sequence[object]) -> StreamingScanKey | None:
return None
def released_stream_as_ended(self, responses_so_far: Sequence[object]) -> tuple[object, ...]:
"""The chunks a client left the stream with, closed the way this endpoint ends a stream, so the
end-of-stream scan also inspects tool calls the stream never finished"""
return tuple(responses_so_far)
def build_block_sse_chunks(
self,
exc: "ModifyResponseException",

View file

@ -17,7 +17,7 @@ This pattern can be replicated for other message formats (e.g., Anthropic).
import json
import time
import uuid
from collections.abc import Mapping, Sequence
from collections.abc import Iterator, Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Union, cast
@ -54,6 +54,7 @@ from litellm.types.utils import (
ChatCompletionDeltaToolCall,
ChatCompletionMessageToolCall,
Choices,
Delta,
GenericGuardrailAPIInputs,
ModelResponse,
ModelResponseStream,
@ -837,6 +838,18 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
tool_calls_in_flight=bool(tool_call_fingerprints) and not stream_ended,
)
def released_stream_as_ended(self, responses_so_far: Sequence[object]) -> tuple[object, ...]:
released_key: Final = self.get_streaming_scan_key(responses_so_far)
if released_key is None or not released_key.tool_calls_in_flight:
return tuple(responses_so_far)
terminator: Final = ModelResponseStream(
choices=[
StreamingChoices(index=index, delta=Delta(), finish_reason="tool_calls")
for index in _choice_indices_with_tool_calls(responses_so_far)
]
)
return (*responses_so_far, terminator)
@staticmethod
def _streamed_tool_call_fingerprints(responses_so_far: Sequence[object]) -> tuple[str, ...]:
return tuple(
@ -1388,6 +1401,20 @@ def _streamed_delta_tool_calls(delta: object) -> tuple[object, ...]:
return stream_item_items(delta, "tool_calls") + legacy
def _released_choices(responses_so_far: Sequence[object]) -> Iterator[object]:
for chunk in responses_so_far:
yield from _stream_chunk_choices(chunk)
def _choice_indices_with_tool_calls(responses_so_far: Sequence[object]) -> tuple[int, ...]:
indices: Final = (
index if isinstance(index := stream_item_field(choice, "index"), int) else 0
for choice in _released_choices(responses_so_far)
if _streamed_delta_tool_calls(stream_item_field(choice, "delta"))
)
return tuple(dict.fromkeys(indices))
def _blocked_stream_identity(
exc: "ModifyResponseException", responses_so_far: Sequence[object]
) -> tuple[str, int, str]:

View file

@ -293,6 +293,51 @@ def _is_tool_call_output_item(item: object) -> bool:
return _tool_call_output_item_mapping(item) is not None
def _released_tool_call_payload(responses_so_far: Sequence[object], item_id: object) -> str | None:
events: Final = tuple(event for event in responses_so_far if stream_item_field(event, "item_id") == item_id)
finished: Final = tuple(
payload
for event in events
if isinstance(event_type := stream_item_field(event, "type"), str)
and event_type in _TOOL_CALL_PAYLOAD_DONE_EVENT_FIELDS
and isinstance(payload := stream_item_field(event, _TOOL_CALL_PAYLOAD_DONE_EVENT_FIELDS[event_type]), str)
)
if finished:
return finished[-1]
deltas: Final = tuple(
delta
for event in events
if stream_item_field(event, "type") in _TOOL_CALL_PAYLOAD_DELTA_EVENT_TYPES
and isinstance(delta := stream_item_field(event, "delta"), str)
)
return "".join(deltas) if deltas else None
def _with_released_payload(item: Mapping[str, object], responses_so_far: Sequence[object]) -> Mapping[str, object]:
field: Final = _TOOL_CALL_PAYLOAD_FIELDS[str(item.get("type"))]
payload: Final = _released_tool_call_payload(responses_so_far, item.get("id"))
return item if payload is None else {**item, field: payload}
def _released_message_item(text: str) -> Mapping[str, object]:
content: Final = [{"type": "output_text", "text": text}]
return {"type": "message", "role": "assistant", "content": content}
def _released_tool_call_items(responses_so_far: Sequence[object]) -> tuple[Mapping[str, object], ...]:
announced: Final = tuple(
item
for item in (
_tool_call_output_item_mapping(stream_item_field(event, "item"))
for event in responses_so_far
if stream_item_field(event, "type") in _OUTPUT_ITEM_EVENT_TYPES
)
if item is not None
)
latest_by_id: Final = MappingProxyType({item.get("id"): item for item in announced})
return tuple(_with_released_payload(item, responses_so_far) for item in latest_by_id.values())
def _last_message_role(messages: Sequence[object]) -> str | None:
if not messages:
return None
@ -1324,6 +1369,27 @@ class OpenAIResponsesHandler(BaseTranslation):
tool_calls_in_flight=self._has_streamed_tool_call_events(responses_so_far),
)
def released_stream_as_ended(self, responses_so_far: Sequence[object]) -> tuple[object, ...]:
if self._check_streaming_has_ended(responses_so_far):
return tuple(responses_so_far)
ends_on_finished_item: Final = (
bool(responses_so_far)
and stream_item_field(responses_so_far[-1], "type") == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE.value
)
if not ends_on_finished_item and not self._has_streamed_tool_call_events(responses_so_far):
return tuple(responses_so_far)
text_events: Final = tuple(
event for event in responses_so_far if stream_item_field(event, "type") in _OUTPUT_TEXT_EVENT_TYPES
)
released_text: Final = self.get_streaming_string_so_far(text_events)
message_items: Final = (_released_message_item(released_text),) if released_text else ()
tool_items: Final = _released_tool_call_items(responses_so_far)
output: Final = [*message_items, *tool_items]
response: Final = {"status": "incomplete", "output": output}
incomplete: Final = ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE.value
envelope: Final = {"type": incomplete, "response": response}
return (*responses_so_far, envelope)
@staticmethod
def _has_streamed_tool_call_events(responses_so_far: Sequence[object]) -> bool:
return any(

View file

@ -272,6 +272,18 @@ def _withheld_provider_output(response: object) -> bool:
return getattr(response, "has_buffered_provider_output", False) is True
async def close_guarded_stream(stream: object) -> None:
if not isinstance(stream, AsyncGenerator):
return
with anyio.CancelScope(shield=True):
try:
await stream.aclose()
except Exception as e: # noqa: BLE001 # a failing callback cleanup must not skip the refund and finalizer
verbose_proxy_logger.warning(
"Closing the guarded stream after a client disconnect raised %s", type(e).__name__
)
def resolve_litellm_call_id(client_call_id: str | None) -> str:
if client_call_id is not None and 0 < len(client_call_id) <= MAX_LITELLM_CALL_ID_LENGTH:
return client_call_id
@ -3909,13 +3921,14 @@ class ProxyBaseLLMRequestProcessing:
client_disconnected = False
delivered_chunk = False
recent_tail = SSE_STREAM_START_TAIL # rebind-ok: rolling window over the yielded bytes
guarded_stream: Final[AsyncGenerator[object, None]] = proxy_logging_obj.async_post_call_streaming_iterator_hook(
user_api_key_dict=user_api_key_dict,
response=response,
request_data=request_data,
)
try:
str_so_far = ""
async for chunk in proxy_logging_obj.async_post_call_streaming_iterator_hook(
user_api_key_dict=user_api_key_dict,
response=response,
request_data=request_data,
):
async for chunk in guarded_stream:
# ``.format(chunk)`` was previously evaluated for every chunk
# regardless of log level; gate it behind the level check.
if debug_enabled:
@ -3971,6 +3984,7 @@ class ProxyBaseLLMRequestProcessing:
# Starlette closes on disconnect, so the nested iterator hook (which
# only sees GeneratorExit on GC) cannot own the refund.
client_disconnected = not stream_completed
await close_guarded_stream(guarded_stream)
if not delivered_chunk and not _withheld_provider_output(response):
from litellm.proxy.spend_tracking.budget_reservation import (
release_budget_reservation_on_cancel,

View file

@ -10,6 +10,7 @@ import sys
sys.path.insert(0, os.path.abspath("../..")) # Adds the parent directory to the system path
import asyncio
import contextlib
import copy
import json
import re
@ -2747,14 +2748,17 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
UnifiedLLMGuardrails,
)
async for streamed_chunk in UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook(
user_api_key_dict=user_api_key_dict,
response=response,
request_data=request_data,
guardrail_to_apply=self,
buffer_until_moderated_default=False,
):
yield streamed_chunk
async with contextlib.aclosing(
UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook(
user_api_key_dict=user_api_key_dict,
response=response,
request_data=request_data,
guardrail_to_apply=self,
buffer_until_moderated_default=False,
)
) as guarded:
async for streamed_chunk in guarded:
yield streamed_chunk
return
# Responses-API events are neither chat-completions chunks nor raw

View file

@ -6,11 +6,14 @@ Unified Guardrail, leveraging LiteLLM's /applyGuardrail endpoint
3. Implements a way to call /applyGuardrail endpoint for `/chat/completions` + `/v1/messages` requests on async_post_call_streaming_iterator_hook
"""
import asyncio
import contextlib
import copy
import json
from collections.abc import AsyncGenerator, AsyncIterable, Awaitable, Callable, Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Protocol
from typing import TYPE_CHECKING, Any, Final, Protocol, TypeAlias, cast
import anyio
from fastapi import HTTPException
from litellm._logging import verbose_proxy_logger
@ -19,6 +22,7 @@ from litellm.cost_calculator import _infer_call_type
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route
from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket
from litellm.llms import get_guardrail_translation_mapping, load_guardrail_translation_mappings
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.guardrails import GuardrailEventHooks
@ -28,6 +32,7 @@ from litellm.types.utils import (
CallTypesLiteral,
Delta,
ModelResponseStream,
StandardLoggingGuardrailInformation,
StreamingChoices,
)
@ -45,6 +50,8 @@ A2A_CALL_TYPES: Final = (CallTypes.asend_message, CallTypes.send_message)
GUARDRAIL_NAME: Final = "unified_llm_guardrails"
_RequestData: TypeAlias = dict[str, object]
class _EndpointTranslation(Protocol):
@property
@ -59,6 +66,9 @@ class _EndpointTranslation(Protocol):
@property
def get_streaming_scan_key(self) -> "Callable[[Sequence[object]], StreamingScanKey | None]": ...
@property
def released_stream_as_ended(self) -> "Callable[[Sequence[object]], tuple[object, ...]]": ...
@property
def build_block_sse_chunks(self) -> "Callable[..., Sequence[bytes] | None]": ...
@ -109,6 +119,18 @@ def _held_choices(held_chars_per_choice: Mapping[int, int]) -> frozenset[int]:
return frozenset(idx for idx, held in held_chars_per_choice.items() if held > 0)
def _recorded_guardrail_information(request_data: _RequestData) -> tuple[StandardLoggingGuardrailInformation, ...]:
_metadata_key, metadata_bucket = get_or_create_metadata_bucket(request_data)
entries: Final = metadata_bucket.get("standard_logging_guardrail_information")
if not isinstance(entries, list):
return ()
return tuple(
cast( # cast-ok: only the guardrail logging helpers write this metadata key
"list[StandardLoggingGuardrailInformation]", entries
)
)
def _is_redundant_scan(scan_key: "StreamingScanKey | None", last_scan_key: "StreamingScanKey | None") -> bool:
if scan_key is None:
return False
@ -601,6 +623,7 @@ class UnifiedLLMGuardrails(CustomLogger):
finish_reason_per_choice: dict[int, str | None],
held_chars_per_choice: dict[int, int],
is_final: bool,
terminated: asyncio.Event,
) -> AsyncGenerator[object, None]:
"""Run one guardrail processing round and emit the resulting diff chunk.
@ -632,6 +655,7 @@ class UnifiedLLMGuardrails(CustomLogger):
is_final=is_final,
)
except ModifyResponseException as e:
terminated.set()
if e.original_response is None:
e.original_response = responses_so_far
async for block_chunk in self.handle_streaming_block(
@ -643,6 +667,7 @@ class UnifiedLLMGuardrails(CustomLogger):
yield block_chunk
raise _StreamTerminated()
except HTTPException as e:
terminated.set()
async for error_item in self.emit_streaming_http_error(
e,
call_type,
@ -664,7 +689,7 @@ class UnifiedLLMGuardrails(CustomLogger):
*,
guardrail_to_apply: CustomGuardrail,
response: AsyncIterable[object],
request_data: dict,
request_data: _RequestData,
user_api_key_dict: UserAPIKeyAuth,
call_type: str,
sampling_rate: int,
@ -687,6 +712,7 @@ class UnifiedLLMGuardrails(CustomLogger):
held_chars_per_choice: Final[dict[int, int]] = {}
chunk_counter = 0
last_chunk: object | None = None
terminated: Final = asyncio.Event()
def _round(reference_chunk: object, is_final: bool) -> AsyncGenerator[object, None]:
return self._emit_transform_round(
@ -702,10 +728,13 @@ class UnifiedLLMGuardrails(CustomLogger):
finish_reason_per_choice=finish_reason_per_choice,
held_chars_per_choice=held_chars_per_choice,
is_final=is_final,
terminated=terminated,
)
saw_tool_calls = False
saw_text_content = False
tool_calls_released = False # rebind-ok: set once a raw tool call reaches the client unscanned
end_of_stream_inspection_started = False # rebind-ok: set once the end-of-stream inspection owns the verdict
try:
async for item in response:
@ -742,6 +771,7 @@ class UnifiedLLMGuardrails(CustomLogger):
held_choices=_held_choices(held_chars_per_choice),
)
responses_yielded.append(tool_only)
tool_calls_released = True
yield tool_only
continue
@ -781,6 +811,7 @@ class UnifiedLLMGuardrails(CustomLogger):
# ``stream_transform_underflow`` 400 from mismatched prefixes. A shallow
# list copy wouldn't help — the mutation is on the chunk objects
# themselves — so we deepcopy.
end_of_stream_inspection_started = True
if saw_tool_calls:
async for out in self._inspect_full_response_for_block(
endpoint_translation=endpoint_translation,
@ -801,6 +832,37 @@ class UnifiedLLMGuardrails(CustomLogger):
yield out
except _StreamTerminated:
return
except (GeneratorExit, asyncio.CancelledError):
await self._scan_uninspected_tool_calls_after_disconnect(
uninspected=tool_calls_released and not end_of_stream_inspection_started and not terminated.is_set(),
endpoint_translation=endpoint_translation,
responses_released=responses_yielded,
guardrail_to_apply=guardrail_to_apply,
user_api_key_dict=user_api_key_dict,
request_data=request_data,
)
raise
@staticmethod
async def _scan_uninspected_tool_calls_after_disconnect(
*,
uninspected: bool,
endpoint_translation: _EndpointTranslation,
responses_released: Sequence[object],
guardrail_to_apply: CustomGuardrail,
user_api_key_dict: UserAPIKeyAuth,
request_data: _RequestData,
) -> None:
if not uninspected:
return
await UnifiedLLMGuardrails._scan_released_stream_after_disconnect(
endpoint_translation=endpoint_translation,
responses_released=responses_released,
last_scan_key=None,
guardrail_to_apply=guardrail_to_apply,
user_api_key_dict=user_api_key_dict,
request_data=request_data,
)
async def _emit_stream_tail(
self,
@ -841,14 +903,15 @@ class UnifiedLLMGuardrails(CustomLogger):
from litellm.integrations.custom_guardrail import ModifyResponseException
try:
await endpoint_translation.process_output_streaming_response(
responses_so_far=responses_so_far,
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=request_data.get("litellm_logging_obj"),
user_api_key_dict=user_api_key_dict,
request_data=request_data,
stream_transform_sink=None,
)
with anyio.CancelScope(shield=bool(responses_yielded)):
await endpoint_translation.process_output_streaming_response(
responses_so_far=responses_so_far,
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=request_data.get("litellm_logging_obj"),
user_api_key_dict=user_api_key_dict,
request_data=request_data,
stream_transform_sink=None,
)
except ModifyResponseException as e:
if e.original_response is None:
e.original_response = responses_so_far
@ -968,6 +1031,45 @@ class UnifiedLLMGuardrails(CustomLogger):
config_value: Final = config.get(name, attribute_value) if isinstance(config, dict) else attribute_value
return self.optional_params.get(name, config_value)
@staticmethod
async def _scan_released_stream_after_disconnect(
*,
endpoint_translation: _EndpointTranslation,
responses_released: Sequence[object],
last_scan_key: "StreamingScanKey | None",
guardrail_to_apply: CustomGuardrail,
user_api_key_dict: UserAPIKeyAuth,
request_data: _RequestData,
) -> None:
scanned: Final = endpoint_translation.released_stream_as_ended(copy.deepcopy(tuple(responses_released)))
if _is_redundant_scan(endpoint_translation.get_streaming_scan_key(scanned), last_scan_key):
return
recorded_before: Final = len(_recorded_guardrail_information(request_data))
with anyio.CancelScope(shield=True):
try:
await endpoint_translation.process_output_streaming_response(
responses_so_far=scanned,
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=request_data.get("litellm_logging_obj"),
user_api_key_dict=user_api_key_dict,
request_data=request_data,
)
except Exception as e: # noqa: BLE001 # the client is gone, so the verdict can only be recorded
verbose_proxy_logger.warning(
"UnifiedLLMGuardrails: %s scanned a stream the client disconnected from and raised %s",
guardrail_to_apply.guardrail_name,
type(e).__name__,
)
recorded_during_scan: Final = _recorded_guardrail_information(request_data)[recorded_before:]
if any(entry.get("guardrail_status") != "success" for entry in recorded_during_scan):
return
guardrail_to_apply.add_standard_logging_guardrail_information_to_request_data(
guardrail_json_response=e,
request_data=request_data,
guardrail_status="guardrail_failed_to_respond",
event_type=GuardrailEventHooks.post_call,
)
async def async_post_call_streaming_iterator_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
@ -995,6 +1097,7 @@ class UnifiedLLMGuardrails(CustomLogger):
if guardrail_to_apply is None:
guardrail_to_apply = request_data.pop("guardrail_to_apply", None)
typed_request_data: Final[_RequestData] = request_data
def _streaming_flag(name: str, default: object) -> Any:
return self.resolve_streaming_flag(guardrail_to_apply, name, default)
@ -1061,17 +1164,20 @@ class UnifiedLLMGuardrails(CustomLogger):
mappings=mappings,
)
if transform_call_type is not None:
async for transformed_item in self._run_incremental_transform_stream(
guardrail_to_apply=guardrail_to_apply,
response=response,
request_data=request_data,
user_api_key_dict=user_api_key_dict,
call_type=transform_call_type,
sampling_rate=sampling_rate,
end_of_stream_only=end_of_stream_only,
mappings=mappings,
):
yield transformed_item
async with contextlib.aclosing(
self._run_incremental_transform_stream(
guardrail_to_apply=guardrail_to_apply,
response=response,
request_data=typed_request_data,
user_api_key_dict=user_api_key_dict,
call_type=transform_call_type,
sampling_rate=sampling_rate,
end_of_stream_only=end_of_stream_only,
mappings=mappings,
)
) as transformed:
async for transformed_item in transformed:
yield transformed_item
return
verbose_proxy_logger.warning(
"UnifiedLLMGuardrails: streaming_transform_mode=incremental_diff is only supported "
@ -1093,217 +1199,240 @@ class UnifiedLLMGuardrails(CustomLogger):
chunks_yielded = False
last_scan_key: StreamingScanKey | None = None # rebind-ok: replaced after every scan round
tool_calls_in_flight = False # rebind-ok: tracks the latest scan key's unscanned tool calls
verdict_settled = False # rebind-ok: set once the end-of-stream scan or a block owns the verdict
async for item in response:
chunk_counter += 1
responses_so_far.append(item)
try:
async for item in response:
chunk_counter += 1
responses_so_far.append(item)
# Infer call type from first chunk if not already done
if call_type is None and user_api_key_dict.request_route is not None:
call_types = get_call_types_for_route(user_api_key_dict.request_route)
if call_types is not None:
call_type = call_types[0].value
# Infer call type from first chunk if not already done
if call_type is None and user_api_key_dict.request_route is not None:
call_types = get_call_types_for_route(user_api_key_dict.request_route)
if call_types is not None:
call_type = call_types[0].value
if call_type is None:
call_type = _infer_call_type(call_type=None, completion_response=item)
if call_type is None:
call_type = _infer_call_type(call_type=None, completion_response=item)
# If call type not supported, just pass through all chunks
if call_type is None or CallTypes(call_type) not in mappings:
yield item
async for remaining_item in response:
yield remaining_item
return
# If call type not supported, just pass through all chunks
if call_type is None or CallTypes(call_type) not in mappings:
yield item
async for remaining_item in response:
yield remaining_item
return
# 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:
if not buffer_until_moderated:
endpoint_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
else:
withheld_items.append(item)
continue
# Process chunk based on sampling rate
if buffer_until_moderated:
withheld_items.append(item)
if chunk_counter % sampling_rate == 0:
endpoint_translation = mappings[CallTypes(call_type)]()
scan_key = endpoint_translation.get_streaming_scan_key(responses_so_far)
if scan_key is not None:
tool_calls_in_flight = scan_key.tool_calls_in_flight
hold_window = buffer_until_moderated and (scan_key is None or tool_calls_in_flight)
if _is_redundant_scan(scan_key, last_scan_key):
verbose_proxy_logger.debug(
"Skipping streaming chunk %s for guardrail %s: nothing new to scan since the last round",
chunk_counter,
guardrail_to_apply.guardrail_name,
)
if buffer_until_moderated:
if hold_window:
continue
for withheld_item in withheld_items:
# 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:
if not buffer_until_moderated:
endpoint_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(withheld_item)
yield withheld_item
withheld_items.clear()
responses_yielded.append(item)
yield item
else:
chunks_yielded = True
responses_yielded.append(item)
yield item
withheld_items.append(item)
continue
# Process chunk based on sampling rate
if buffer_until_moderated:
withheld_items.append(item)
if chunk_counter % sampling_rate == 0:
endpoint_translation = mappings[CallTypes(call_type)]()
scan_key = endpoint_translation.get_streaming_scan_key(responses_so_far)
if scan_key is not None:
tool_calls_in_flight = scan_key.tool_calls_in_flight
hold_window = buffer_until_moderated and (scan_key is None or tool_calls_in_flight)
if _is_redundant_scan(scan_key, last_scan_key):
verbose_proxy_logger.debug(
"Skipping streaming chunk %s for guardrail %s: nothing new to scan since the last round",
chunk_counter,
guardrail_to_apply.guardrail_name,
)
if buffer_until_moderated:
if hold_window:
continue
for withheld_item in withheld_items:
chunks_yielded = True
responses_yielded.append(withheld_item)
yield withheld_item
withheld_items.clear()
else:
chunks_yielded = True
responses_yielded.append(item)
yield item
continue
verbose_proxy_logger.debug(
"Processing streaming chunk %s (sampling_rate=%s) with guardrail %s",
chunk_counter,
sampling_rate,
guardrail_to_apply.guardrail_name,
)
original_items = (
tuple(copy.deepcopy(withheld_items)) if buffer_until_moderated else (copy.deepcopy(item),)
)
try:
await endpoint_translation.process_output_streaming_response(
responses_so_far=responses_so_far,
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=request_data.get("litellm_logging_obj"),
user_api_key_dict=user_api_key_dict,
request_data=request_data,
)
except ModifyResponseException as e:
verdict_settled = True
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:
verdict_settled = True
# Response already started (we already yielded chunks); cannot send 400.
async for error_item in self.emit_streaming_http_error(
e,
call_type,
responses_so_far,
request_data,
endpoint_translation=endpoint_translation,
stream_started=chunks_yielded,
responses_yielded=responses_yielded,
):
yield error_item
return
if scan_key is not None:
last_scan_key = scan_key
if hold_window:
verbose_proxy_logger.debug(
"Holding %s buffered chunks for guardrail %s: this round could not scan the whole window",
len(withheld_items),
guardrail_to_apply.guardrail_name,
)
withheld_items[:] = original_items
continue
for original_item in original_items:
chunks_yielded = True
responses_yielded.append(original_item)
yield original_item
withheld_items.clear()
else:
if not buffer_until_moderated:
chunks_yielded = True
responses_yielded.append(item)
yield item
# Stream has ended - do final processing with all collected chunks
if call_type is not None and CallTypes(call_type) in mappings:
verbose_proxy_logger.debug(
"Processing streaming chunk %s (sampling_rate=%s) with guardrail %s",
chunk_counter,
sampling_rate,
"Processing final streaming response with all %s chunks for guardrail %s",
len(responses_so_far),
guardrail_to_apply.guardrail_name,
)
original_items = (
tuple(copy.deepcopy(withheld_items)) if buffer_until_moderated else (copy.deepcopy(item),)
endpoint_translation = mappings[CallTypes(call_type)]()
buffered_items: Final = (
tuple(copy.deepcopy(withheld_items))
if buffer_until_moderated and release_on_scan and not end_of_stream_only
else tuple(withheld_items)
if buffer_until_moderated
else None
)
end_scan_key: Final = endpoint_translation.get_streaming_scan_key(responses_so_far)
verdict_settled = True
if _is_redundant_scan(end_scan_key, last_scan_key):
verbose_proxy_logger.debug(
"Skipping end-of-stream scan for guardrail %s: the last sampled round already scanned it all",
guardrail_to_apply.guardrail_name,
)
for buffered_item in buffered_items or ():
yield buffered_item
for pending_item in pending_end_of_stream_items:
responses_yielded.append(pending_item)
yield pending_item
return
try:
await endpoint_translation.process_output_streaming_response(
responses_so_far=responses_so_far,
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=request_data.get("litellm_logging_obj"),
user_api_key_dict=user_api_key_dict,
request_data=request_data,
)
with anyio.CancelScope(shield=chunks_yielded):
await endpoint_translation.process_output_streaming_response(
responses_so_far=responses_so_far,
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=request_data.get("litellm_logging_obj"),
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
# 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.
# 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=chunks_yielded,
stream_started=bool(responses_yielded),
responses_so_far=responses_yielded,
):
yield block_chunk
return
except HTTPException as e:
# Response already started (we already yielded chunks); cannot send 400.
async for error_item in self.emit_streaming_http_error(
e,
call_type,
responses_so_far,
request_data,
endpoint_translation=endpoint_translation,
stream_started=chunks_yielded,
stream_started=bool(responses_yielded),
responses_yielded=responses_yielded,
):
yield error_item
return
if scan_key is not None:
last_scan_key = scan_key
if hold_window:
verbose_proxy_logger.debug(
"Holding %s buffered chunks for guardrail %s: this round could not scan the whole window",
len(withheld_items),
guardrail_to_apply.guardrail_name,
)
withheld_items[:] = original_items
continue
for original_item in original_items:
chunks_yielded = True
responses_yielded.append(original_item)
yield original_item
withheld_items.clear()
else:
if not buffer_until_moderated:
chunks_yielded = True
responses_yielded.append(item)
yield item
# Stream has ended - do final processing with all collected chunks
if call_type is not None and CallTypes(call_type) in mappings:
verbose_proxy_logger.debug(
"Processing final streaming response with all %s chunks for guardrail %s",
len(responses_so_far),
guardrail_to_apply.guardrail_name,
)
endpoint_translation = mappings[CallTypes(call_type)]()
buffered_items: Final = (
tuple(copy.deepcopy(withheld_items))
if buffer_until_moderated and release_on_scan and not end_of_stream_only
else tuple(withheld_items)
if buffer_until_moderated
else None
)
end_scan_key: Final = endpoint_translation.get_streaming_scan_key(responses_so_far)
if _is_redundant_scan(end_scan_key, last_scan_key):
verbose_proxy_logger.debug(
"Skipping end-of-stream scan for guardrail %s: the last sampled round already scanned it all",
guardrail_to_apply.guardrail_name,
)
for buffered_item in buffered_items or ():
yield buffered_item
for pending_item in pending_end_of_stream_items:
responses_yielded.append(pending_item)
yield pending_item
return
try:
await endpoint_translation.process_output_streaming_response(
responses_so_far=responses_so_far,
except (GeneratorExit, asyncio.CancelledError):
translation_class: Final = None if call_type is None else mappings.get(CallTypes(call_type))
if (
chunks_yielded
and not verdict_settled
and translation_class is not None
and isinstance(guardrail_to_apply, CustomGuardrail)
):
await self._scan_released_stream_after_disconnect(
endpoint_translation=translation_class(),
responses_released=responses_yielded,
last_scan_key=last_scan_key,
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=request_data.get("litellm_logging_obj"),
user_api_key_dict=user_api_key_dict,
request_data=request_data,
request_data=typed_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:
async for error_item in self.emit_streaming_http_error(
e,
call_type,
responses_so_far,
request_data,
endpoint_translation=endpoint_translation,
stream_started=bool(responses_yielded),
responses_yielded=responses_yielded,
):
yield error_item
raise

View file

@ -399,6 +399,7 @@ from litellm.proxy.common_request_processing import (
ProxyBaseLLMRequestProcessing,
_is_azure_model_router_request,
_should_return_raw_model_name,
close_guarded_stream,
create_response,
log_llm_api_exception,
open_sse_before_first_byte,
@ -9911,6 +9912,17 @@ async def async_data_generator(
stream_completed = False
client_disconnected = False
error_state: Final = ResponsesStreamErrorState() if responses_stream_errors else None
needs_iterator_wrap: Final = proxy_logging_obj.needs_iterator_wrap()
stream_iterator: Final[AsyncIterator[object]] = (
proxy_logging_obj.async_post_call_streaming_iterator_hook(
user_api_key_dict=user_api_key_dict,
response=response,
request_data=request_data,
)
if needs_iterator_wrap
else response
)
stream_source: AsyncIterator[object] | None = None # rebind-ok: bound once the keepalive policy resolves
try:
error_message: str | None = None
requested_model_from_client: Final = _get_client_requested_model_for_streaming(request_data=request_data)
@ -9935,21 +9947,11 @@ async def async_data_generator(
# per-chunk hook. Coalescing them into a single flag forced wasted
# ``get_response_string`` work per chunk on every deployment that
# happened to ship a streaming-iterator override (the default).
needs_iterator_wrap: Final = proxy_logging_obj.needs_iterator_wrap()
needs_per_chunk_hook: Final = proxy_logging_obj.needs_per_chunk_streaming_hook()
is_raw_sse_stream: Final = bool(request_data.get("_litellm_raw_sse_stream"))
strip_stream_usage: Final = bool(request_data.get("_litellm_strip_stream_usage"))
raw_sse_buffer = ""
if needs_iterator_wrap:
stream_iterator = proxy_logging_obj.async_post_call_streaming_iterator_hook(
user_api_key_dict=user_api_key_dict,
response=response,
request_data=request_data,
)
else:
stream_iterator = response
# A stream can start on a deployment with keepalive off and fall back
# mid-stream to one that enables it: only skip wrapping altogether when
# there's no router to ever fall back through AND the resolved interval
@ -9958,7 +9960,7 @@ async def async_data_generator(
# happens to start with it off.
resolve_keepalive_seconds: Final = _make_keepalive_resolver(request_data)
initial_keepalive_seconds: Final = resolve_keepalive_seconds(response)
stream_source: Final = (
stream_source = (
_iter_with_keepalive(
stream_iterator.__aiter__(),
resolve_keepalive_seconds,
@ -10099,6 +10101,9 @@ async def async_data_generator(
# (a nested iterator hook would only see GeneratorExit on GC).
if not stream_completed:
client_disconnected = True
for guarded_layer in (stream_source, stream_iterator):
if guarded_layer is not response:
await close_guarded_stream(guarded_layer)
raise
except Exception as e:
verbose_proxy_logger.exception("litellm.proxy.proxy_server.async_data_generator(): Exception occured - %s", e)

View file

@ -44,6 +44,7 @@ from typing import (
Union,
cast,
overload,
runtime_checkable,
)
from typing_extensions import ReadOnly, TypedDict
@ -524,8 +525,13 @@ class _UpstreamStreamBoundary(Generic[_T]):
raise
@runtime_checkable
class _ClosableAsyncIterator(Protocol):
def aclose(self) -> object: ...
class _StreamIteratorHook(Protocol[_T]):
def __call__(self, *, response: AsyncIterator[_T]) -> AsyncGenerator[_T, None]: ...
def __call__(self, *, response: AsyncIterator[_T]) -> AsyncIterator[_T]: ...
def _is_client_error_exception(exc: Exception) -> bool:
@ -2773,8 +2779,22 @@ class ProxyLogging:
) -> AsyncGenerator[_T, None]:
upstream: Final = _UpstreamStreamBoundary(response)
try:
async for chunk in hook(response=upstream):
yield chunk
guarded: Final = hook(response=upstream)
try:
async for chunk in guarded:
yield chunk
finally:
if isinstance(guarded, _ClosableAsyncIterator):
try:
closing: Final = guarded.aclose()
if inspect.isawaitable(closing):
await closing
except Exception as e: # noqa: BLE001 # a finished stream must not fail on callback cleanup
verbose_proxy_logger.warning(
"Closing the streaming iterator of %s raised %s",
getattr(callback, "guardrail_name", None) or type(callback).__name__,
type(e).__name__,
)
except Exception as e:
if e is not upstream.failure:
enrich_http_exception_with_guardrail_context(e, callback)
@ -3959,6 +3979,7 @@ class ProxyLogging:
stream_needs_translation: Final = ProxyLogging._stream_requires_guardrail_translation(user_api_key_dict)
pipeline_gated_names: Final = _pipeline_step_guardrail_names(post_call_pipelines)
guarded_layers: Final[list[AsyncGenerator[object, None]]] = [] # mutable-ok: closed on disconnect
for resolved_callback, kind in caps.iterator_overrides:
if isinstance(resolved_callback, CustomGuardrail):
if resolved_callback.guardrail_name in pipeline_gated_names:
@ -4001,6 +4022,7 @@ class ProxyLogging:
hook,
request_data=request_data,
)
guarded_layers.append(current_response)
pipeline_translation: Final = (
resolve_endpoint_translation(user_api_key_dict, None) if post_call_pipelines else None
@ -4013,6 +4035,7 @@ class ProxyLogging:
pipelines=post_call_pipelines,
translation=pipeline_translation,
)
guarded_layers.append(current_response)
served_chunks: Final[list[object]] = [] # mutable-ok: accumulates while yielding to the client
try:
@ -4020,6 +4043,7 @@ class ProxyLogging:
served_chunks.append(chunk)
yield chunk
except (GeneratorExit, asyncio.CancelledError):
await ProxyLogging._close_guarded_layers(guarded_layers)
ProxyLogging._record_served_stream_output(request_data, served_chunks)
raise
except Exception as e:
@ -4100,6 +4124,16 @@ class ProxyLogging:
for buffered_item in buffered:
yield buffered_item
@staticmethod
async def _close_guarded_layers(layers: Sequence[AsyncGenerator[object, None]]) -> None:
for layer in reversed(layers):
try:
await layer.aclose()
except Exception as e: # noqa: BLE001 # one failing callback cleanup must not skip the inner ones
verbose_proxy_logger.warning(
"Closing a streaming callback layer after a client disconnect raised %s", type(e).__name__
)
@staticmethod
def _record_served_stream_output(request_data: Mapping[str, object], served_chunks: Sequence[object]) -> None:
logging_obj: Final = request_data.get("litellm_logging_obj")

View file

@ -1,12 +1,20 @@
import asyncio
import json
import os
import re
import signal
import socket
import threading
import uuid
from collections.abc import Callable, Iterator, Mapping
from concurrent.futures import ThreadPoolExecutor
from contextlib import contextmanager
from dataclasses import dataclass
from pathlib import Path
from types import MappingProxyType
from typing import Final
import anthropic
import httpx
import psutil
import pytest
@ -14,9 +22,10 @@ import yaml
from integration._support.client import Gateway, eventually, object_value
from integration._support.database import read_rows
from integration._support.mcp import mcp_peer, register_mcp, tool_names
from integration._support.process import group_members, owned_proxy, owned_proxy_process
from integration._support.wire import Reply, Request, wire_server
from integration._support.process import OwnedProxy, group_members, owned_proxy, owned_proxy_process
from integration._support.wire import Reply, Request, Wire, wire_server
from openai import AsyncOpenAI, OpenAI
from pydantic import JsonValue
@pytest.mark.covers("other.observability.guardrails.rewrite_reaches_correct_anthropic_positions")
@ -1796,3 +1805,739 @@ def test_responses_pre_call_denial_stream_survives_worker_kill(gateway: Gateway,
for response in responses:
assert response.status_code == 200, response.text
assert response.headers["content-type"].startswith("text/event-stream"), response.text
_TOKEN: Final = re.compile(rb"token-[0-9a-f]{32}-\d+")
def _secret_for(request: Request) -> str:
token: Final = _TOKEN.search(request.body)
assert token is not None, request.body
return "synthetic-leaked-secret-" + token.group().decode()
def _chat_frame(identity: str, choices: tuple[dict[str, JsonValue], ...]) -> bytes:
payload: Final = {
"id": identity,
"object": "chat.completion.chunk",
"created": 1,
"model": "gpt-4o-mini",
"choices": list(choices),
}
return b"data: " + json.dumps(payload).encode() + b"\n\n"
def _chat_choice(index: int, delta: dict[str, JsonValue], finish: str | None = None) -> dict[str, JsonValue]:
return {"index": index, "delta": delta, "finish_reason": finish}
def _chat_stream_frames(secret: str, shape: str) -> tuple[bytes, ...]:
identity: Final = "chatcmpl-" + uuid.uuid4().hex
tool_call: Final = {
"index": 0,
"id": "call_" + identity,
"type": "function",
"function": {"name": "lookup", "arguments": json.dumps({"query": secret})},
}
released: Final = {
"text": (_chat_choice(0, {"role": "assistant", "content": secret}),),
"empty": (_chat_choice(0, {"role": "assistant", "content": ""}),),
"tool_call": (_chat_choice(0, {"role": "assistant", "tool_calls": [tool_call]}),),
"two_choices": (
_chat_choice(0, {"role": "assistant", "content": secret + "-first"}),
_chat_choice(1, {"role": "assistant", "content": secret + "-second"}),
),
}[shape]
finish: Final = "tool_calls" if shape == "tool_call" else "stop"
tail: Final = tuple(_chat_choice(int(str(choice["index"])), {"content": " tail"}, finish) for choice in released)
return (_chat_frame(identity, released), _chat_frame(identity, tail), b"data: [DONE]\n\n")
def _chat_completion_body(secret: str) -> bytes:
return json.dumps(
{
"id": "chatcmpl-" + uuid.uuid4().hex,
"object": "chat.completion",
"created": 1,
"model": "gpt-4o-mini",
"choices": [{"index": 0, "message": {"role": "assistant", "content": secret}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 9, "completion_tokens": 4, "total_tokens": 13},
}
).encode()
def _responses_tool_call_frames(secret: str, *, with_text: bool) -> tuple[bytes, ...]:
identity: Final = "resp_" + uuid.uuid4().hex
arguments: Final = json.dumps({"query": secret})
pending: Final = {"type": "function_call", "id": "fc_" + identity, "call_id": "call_" + identity, "name": "lookup"}
finished: Final = {**pending, "arguments": arguments, "status": "completed"}
envelope: Final = {"id": identity, "object": "response", "created_at": 1, "model": "gpt-4o-mini", "output": []}
message: Final = {
"type": "message",
"id": "msg_" + identity,
"role": "assistant",
"status": "completed",
"content": [{"type": "output_text", "text": secret, "annotations": []}],
}
text_delta: Final = {
"type": "response.output_text.delta",
"item_id": "msg_" + identity,
"output_index": 0,
"content_index": 0,
"delta": secret,
}
text_events: Final = (text_delta,) if with_text else ()
tool_index: Final = len(text_events)
output: Final = [*((message,) if with_text else ()), finished]
events: Final = (
{"type": "response.created", "response": {**envelope, "status": "in_progress"}},
*text_events,
{
"type": "response.output_item.added",
"output_index": tool_index,
"item": {**pending, "arguments": "", "status": "in_progress"},
},
{
"type": "response.function_call_arguments.delta",
"item_id": "fc_" + identity,
"output_index": tool_index,
"delta": arguments,
},
{"type": "response.output_item.done", "output_index": tool_index, "item": finished},
{"type": "response.completed", "response": {**envelope, "status": "completed", "output": output}},
)
encoded: Final = tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events)
released_through: Final = len(events) - 1 if with_text else 3
return (b"".join(encoded[:released_through]), b"".join(encoded[released_through:]))
def _responses_stream_frames(secret: str) -> tuple[bytes, ...]:
identity: Final = "resp_" + uuid.uuid4().hex
completed: Final = {
"id": identity,
"object": "response",
"created_at": 1,
"status": "completed",
"model": "gpt-4o-mini",
"output": [
{
"type": "message",
"id": "msg_" + identity,
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": secret, "annotations": []}],
}
],
"usage": {
"input_tokens": 11,
"output_tokens": 4,
"total_tokens": 15,
"input_tokens_details": {"cached_tokens": 0},
"output_tokens_details": {"reasoning_tokens": 0},
},
}
events: Final = (
{"type": "response.created", "response": {**completed, "status": "in_progress", "output": [], "usage": None}},
{
"type": "response.output_text.delta",
"item_id": "msg_" + identity,
"output_index": 0,
"content_index": 0,
"delta": secret,
},
{"type": "response.completed", "response": completed},
)
encoded: Final = tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events)
return (encoded[0] + encoded[1], encoded[2])
def _gemini_stream_frames(secret: str) -> tuple[bytes, ...]:
def frame(text: str, finish: str | None) -> bytes:
candidate: Final = {
"content": {"parts": [{"text": text}], "role": "model"},
"index": 0,
**({"finishReason": finish} if finish else {}),
}
payload: Final = {
"candidates": [candidate],
"usageMetadata": {"promptTokenCount": 10, "candidatesTokenCount": 5, "totalTokenCount": 15},
"modelVersion": "gemini-2.5-flash",
}
return b"data: " + json.dumps(payload).encode() + b"\r\n\r\n"
return (frame(secret, None), frame(" tail", "STOP"))
def _scripted_provider(gate: threading.Event | None, pause: float, shape: str) -> Callable[[Request], Reply]:
def provider(request: Request) -> Reply:
secret: Final = _secret_for(request)
path: Final = request.target.split("?")[0]
if path.endswith("/chat/completions") and not json.loads(request.body).get("stream"):
return Reply(body=_chat_completion_body(secret))
frames: Final = (
_gemini_stream_frames(secret)
if "streamGenerateContent" in path
else (
_responses_tool_call_frames(secret, with_text=shape == "text_then_tool_call")
if shape in ("tool_call", "text_then_tool_call")
else _responses_stream_frames(secret)
)
if path.endswith("/responses")
else _chat_stream_frames(secret, shape)
)
return Reply(content_type="text/event-stream", chunks=frames, gate_after_first=gate, pause_between_chunks=pause)
return provider
def _allowing_guardrail(request: Request) -> Reply:
assert request.target == "/beta/litellm_basic_guardrail_api", request.target
return Reply(body=json.dumps({"action": "NONE"}).encode())
def _failing_response_scans(reply: Reply) -> Callable[[Request], Reply]:
def guardrail(request: Request) -> Reply:
assert request.target == "/beta/litellm_basic_guardrail_api", request.target
if json.loads(request.body)["input_type"] == "response":
return reply
return Reply(body=json.dumps({"action": "NONE"}).encode())
return guardrail
def _post_call_config(
tmp_path: Path, identity: str, policy_url: str, params: Mapping[str, JsonValue], default_on: bool
) -> Path:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["guardrails"] = [
{
"guardrail_name": identity,
"litellm_params": {
"guardrail": "generic_guardrail_api",
"mode": "post_call",
"default_on": default_on,
"api_base": policy_url,
"api_key": "synthetic-guardrail-key",
**params,
},
}
]
path: Final = tmp_path / f"{identity}.yaml"
path.write_text(yaml.safe_dump(config))
return path
class _ScanLog:
def __init__(self, policy: Wire) -> None:
self.policy: Final = policy
self.seen: tuple[dict[str, JsonValue], ...] = ()
def response_scans(self, secret: str) -> tuple[dict[str, JsonValue], ...]:
self.seen = (*self.seen, *(object_value(json.loads(request.body)) for request in self.policy.drain()))
return tuple(body for body in self.seen if body["input_type"] == "response" and secret in json.dumps(body))
@dataclass(frozen=True, slots=True)
class _DisconnectRig:
owned: OwnedProxy
model: str
gemini: str
scans: _ScanLog
identity: str
gate: threading.Event
upstream: Wire
@property
def candidate(self) -> Gateway:
return self.owned.gateway
def token(self, index: int = 0) -> str:
return f"token-{self.identity.removeprefix('guardrail')}-{index}"
def secret(self, index: int = 0) -> str:
return "synthetic-leaked-secret-" + self.token(index)
_END_OF_STREAM_ONLY: Final = MappingProxyType({"streaming_end_of_stream_only": True})
@contextmanager
def _disconnect_rig(
gateway: Gateway,
tmp_path: Path,
*,
params: Mapping[str, JsonValue] = _END_OF_STREAM_ONLY,
guardrail: Callable[[Request], Reply] = _allowing_guardrail,
gated: bool = True,
pause: float = 0,
shape: str = "text",
default_on: bool = True,
workers: int = 1,
) -> Iterator[_DisconnectRig]:
identity: Final = "guardrail" + uuid.uuid4().hex
gate: Final = threading.Event()
with (
wire_server(guardrail) as policy,
wire_server(_scripted_provider(gate if gated else None, pause, shape)) as upstream,
):
config: Final = _post_call_config(tmp_path, identity, policy.url, params, default_on)
try:
with (
owned_proxy_process(gateway, tmp_path, {}, config=config, workers=workers) as owned,
owned.gateway.scenario() as scenario,
):
yield _DisconnectRig(
owned,
scenario.model(api_base=upstream.url + "/v1", api_key="synthetic-openai-key"),
scenario.model(
model="gemini/gemini-2.5-flash", api_base=upstream.url, api_key="synthetic-gemini-key"
),
_ScanLog(policy),
identity,
gate,
upstream,
)
finally:
gate.set()
def _close_on(response: httpx.Response, marker: str) -> str:
assert response.status_code == 200, response.read()
for line in response.iter_lines():
if marker in line:
return line
raise AssertionError(f"The stream ended before the client received {marker}")
def _stream_and_close(rig: _DisconnectRig, path: str, body: Mapping[str, JsonValue], marker: str) -> str:
with rig.candidate.client.stream(
"POST", path, json=dict(body), headers={"Authorization": f"Bearer {rig.candidate.key}"}
) as response:
return _close_on(response, marker)
def _chat_body(rig: _DisconnectRig, index: int, stream: bool = True) -> dict[str, JsonValue]:
return {
"model": rig.model,
"messages": [{"role": "user", "content": "synthetic prompt " + rig.token(index)}],
"stream": stream,
}
def _chat_httpx(rig: _DisconnectRig, index: int = 0) -> str:
return _stream_and_close(rig, "/v1/chat/completions", _chat_body(rig, index), rig.secret(index))
def _responses_httpx(rig: _DisconnectRig, index: int = 0) -> str:
body: Final = {"model": rig.model, "input": "synthetic prompt " + rig.token(index), "stream": True}
return _stream_and_close(rig, "/v1/responses", body, rig.secret(index))
def _messages_httpx(rig: _DisconnectRig, index: int = 0) -> str:
body: Final = {**_chat_body(rig, index), "max_tokens": 64}
return _stream_and_close(rig, "/v1/messages", body, rig.secret(index))
def _gemini_httpx(rig: _DisconnectRig, index: int = 0) -> str:
body: Final = {"contents": [{"role": "user", "parts": [{"text": "synthetic prompt " + rig.token(index)}]}]}
path: Final = f"/v1beta/models/{rig.gemini}:streamGenerateContent?alt=sse"
return _stream_and_close(rig, path, body, rig.secret(index))
def _chat_async_openai_sdk(rig: _DisconnectRig, index: int = 0) -> str:
async def read() -> str:
client: Final = AsyncOpenAI(
base_url=str(rig.candidate.client.base_url) + "/v1",
api_key=rig.candidate.key,
max_retries=0,
http_client=httpx.AsyncClient(trust_env=False, timeout=30),
)
async with client:
stream: Final = await client.chat.completions.create(
model=rig.model,
messages=[{"role": "user", "content": "synthetic prompt " + rig.token(index)}],
stream=True,
)
async for chunk in stream:
if chunk.choices and rig.secret(index) in (chunk.choices[0].delta.content or ""):
await stream.close()
return chunk.choices[0].delta.content or ""
raise AssertionError("The stream ended before the client received the streamed content")
return asyncio.run(read())
def _responses_openai_sdk(rig: _DisconnectRig, index: int = 0) -> str:
client: Final = OpenAI(
base_url=str(rig.candidate.client.base_url) + "/v1",
api_key=rig.candidate.key,
max_retries=0,
http_client=httpx.Client(trust_env=False, timeout=30),
)
with client:
stream: Final = client.responses.create(
model=rig.model, input="synthetic prompt " + rig.token(index), stream=True
)
for event in stream:
if event.type == "response.output_text.delta" and rig.secret(index) in event.delta:
stream.close()
return event.delta
raise AssertionError("The stream ended before the client received the streamed content")
def _messages_anthropic_sdk(rig: _DisconnectRig, index: int = 0) -> str:
client: Final = anthropic.Anthropic(
base_url=str(rig.candidate.client.base_url),
api_key=rig.candidate.key,
max_retries=0,
http_client=httpx.Client(trust_env=False, timeout=30),
)
with client:
stream: Final = client.messages.create(
model=rig.model,
max_tokens=64,
messages=[{"role": "user", "content": "synthetic prompt " + rig.token(index)}],
stream=True,
)
for event in stream:
text: Final = (
event.delta.text if event.type == "content_block_delta" and event.delta.type == "text_delta" else ""
)
if rig.secret(index) in text:
stream.close()
return text
raise AssertionError("The stream ended before the client received the streamed content")
def _scanned_while_upstream_is_held(
rig: _DisconnectRig, disconnect: Callable[[_DisconnectRig, int], str], index: int = 0
) -> tuple[dict[str, JsonValue], ...]:
try:
received: Final = disconnect(rig, index)
assert rig.secret(index) in received, received
return eventually(
lambda: rig.scans.response_scans(rig.secret(index)), lambda values: len(values) >= 1, seconds=4
)
finally:
rig.gate.set()
def _post_call_statuses(rig: _DisconnectRig, model: str, rows: int = 1) -> tuple[tuple[str, ...], ...]:
found: Final = eventually(
lambda: read_rows('SELECT metadata FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)),
lambda values: len(values) == rows,
seconds=70,
)
def statuses(metadata: JsonValue) -> tuple[str, ...]:
entries: Final = object_value(metadata).get("guardrail_information") or []
assert isinstance(entries, list), metadata
post_call: Final = tuple(
entry
for entry in (object_value(value) for value in entries)
if entry.get("guardrail_name") == rig.identity and entry.get("guardrail_mode") == "post_call"
)
return tuple(str(entry["guardrail_status"]) for entry in post_call)
return tuple(statuses(row["metadata"]) for row in found)
_ENDPOINT_CLIENTS: Final = (
pytest.param(_chat_httpx, id="chat-httpx"),
pytest.param(_chat_async_openai_sdk, id="chat-async-openai-sdk"),
pytest.param(_responses_httpx, id="responses-httpx"),
pytest.param(_responses_openai_sdk, id="responses-openai-sdk"),
pytest.param(_messages_httpx, id="messages-httpx"),
pytest.param(_messages_anthropic_sdk, id="messages-anthropic-sdk"),
pytest.param(_gemini_httpx, id="native-gemini-stream-generate-content"),
)
_NO_DISCONNECT_ROW_LIT_8603: Final = pytest.mark.skip(
reason="BUG: LIT-8603 a mid-stream disconnect writes no spend row"
)
@pytest.mark.parametrize("disconnect", _ENDPOINT_CLIENTS)
def test_client_disconnect_mid_stream_still_scans_the_content_it_already_received(
gateway: Gateway, tmp_path: Path, disconnect: Callable[[_DisconnectRig, int], str]
) -> None:
with _disconnect_rig(gateway, tmp_path) as rig:
scans: Final = _scanned_while_upstream_is_held(rig, disconnect)
assert len(scans) == 1, scans
@pytest.mark.parametrize(
"disconnect",
(
pytest.param(_chat_httpx, id="chat-httpx"),
pytest.param(_chat_async_openai_sdk, id="chat-async-openai-sdk"),
pytest.param(_responses_httpx, id="responses-httpx", marks=_NO_DISCONNECT_ROW_LIT_8603),
pytest.param(_messages_anthropic_sdk, id="messages-anthropic-sdk", marks=_NO_DISCONNECT_ROW_LIT_8603),
pytest.param(
_gemini_httpx,
id="native-gemini-stream-generate-content",
marks=pytest.mark.skip(reason="BUG: LIT-9087 a mid-stream disconnect writes no spend row"),
),
),
)
def test_client_disconnect_mid_stream_records_the_post_call_verdict_on_the_spend_row(
gateway: Gateway, tmp_path: Path, disconnect: Callable[[_DisconnectRig, int], str]
) -> None:
with _disconnect_rig(gateway, tmp_path) as rig:
_scanned_while_upstream_is_held(rig, disconnect)
model: Final = rig.gemini if disconnect is _gemini_httpx else rig.model
assert _post_call_statuses(rig, model) == (("success",),)
@pytest.mark.parametrize(
"reply",
(
pytest.param(Reply(status=500, body=b'{"error": "synthetic guardrail outage"}'), id="guardrail-500"),
pytest.param(Reply(body=b"synthetic non-json guardrail body"), id="guardrail-malformed-200"),
),
)
def test_client_disconnect_mid_stream_records_a_failed_scan_when_the_guardrail_errors(
gateway: Gateway, tmp_path: Path, reply: Reply
) -> None:
with _disconnect_rig(gateway, tmp_path, guardrail=_failing_response_scans(reply)) as rig:
_scanned_while_upstream_is_held(rig, _chat_httpx)
assert _post_call_statuses(rig, rig.model) == (("guardrail_failed_to_respond",),)
def test_client_disconnect_mid_stream_records_a_blocking_verdict_and_keeps_serving(
gateway: Gateway, tmp_path: Path
) -> None:
blocked: Final = Reply(body=json.dumps({"action": "BLOCKED", "blocked_reason": "synthetic leak"}).encode())
with _disconnect_rig(gateway, tmp_path, guardrail=_failing_response_scans(blocked)) as rig:
_scanned_while_upstream_is_held(rig, _chat_httpx)
statuses: Final = _post_call_statuses(rig, rig.model)
assert len(statuses) == 1 and len(statuses[0]) == 1 and statuses[0][0] != "success", statuses
health: Final = rig.candidate.request("GET", "/health/liveliness")
assert health.status_code == 200, health.text
def test_client_disconnect_while_end_of_stream_scan_is_in_flight_still_records_the_verdict(
gateway: Gateway, tmp_path: Path
) -> None:
scan_started: Final = threading.Event()
scan_released: Final = threading.Event()
def guardrail(request: Request) -> Reply:
assert request.target == "/beta/litellm_basic_guardrail_api", request.target
if json.loads(request.body)["input_type"] == "response":
scan_started.set()
assert scan_released.wait(timeout=30), "The in-flight scan was never released"
return Reply(body=json.dumps({"action": "NONE"}).encode())
with _disconnect_rig(gateway, tmp_path, guardrail=guardrail, gated=False) as rig:
with rig.candidate.client.stream(
"POST",
"/v1/chat/completions",
json=_chat_body(rig, 0),
headers={"Authorization": f"Bearer {rig.candidate.key}"},
) as response:
try:
assert rig.secret() in _close_on(response, rig.secret())
assert scan_started.wait(timeout=10), "The end-of-stream scan never started"
finally:
pass
scan_released.set()
assert _post_call_statuses(rig, rig.model) == (("success",),)
assert len(rig.scans.response_scans(rig.secret())) == 1
@pytest.mark.parametrize(
("params", "shape", "expected"),
(
pytest.param(
{"streaming_buffer_until_moderated": False},
"text",
("synthetic-leaked-secret-",),
id="sampled-before-the-sampling-threshold",
),
pytest.param(
{"streaming_buffer_until_moderated": False, "streaming_transform_mode": "incremental_diff"},
"tool_call",
('\\"query\\": \\"synthetic-leaked-secret-',),
id="incremental-diff-tool-call-in-flight",
),
pytest.param(dict(_END_OF_STREAM_ONLY), "two_choices", ("-first", "-second"), id="two-choices"),
),
)
def test_client_disconnect_mid_stream_scans_what_each_streaming_mode_released(
gateway: Gateway, tmp_path: Path, params: dict[str, JsonValue], shape: str, expected: tuple[str, ...]
) -> None:
with _disconnect_rig(gateway, tmp_path, params=params, shape=shape) as rig:
marker: Final = rig.secret() + ("-first" if shape == "two_choices" else "")
try:
received: Final = _stream_and_close(rig, "/v1/chat/completions", _chat_body(rig, 0), rig.token())
assert rig.token() in received, received
scans: Final = eventually(
lambda: rig.scans.response_scans(rig.secret()), lambda values: len(values) >= 1, seconds=4
)
finally:
rig.gate.set()
payload: Final = json.dumps(scans[-1])
assert all(fragment in payload for fragment in expected), (marker, scans)
assert _post_call_statuses(rig, rig.model)[0][-1:] == ("success",)
def _tool_call_request(rig: _DisconnectRig, path: str) -> dict[str, JsonValue]:
if path == "/v1/responses":
return {"model": rig.model, "input": "synthetic prompt " + rig.token(), "stream": True}
if path == "/v1/messages":
return {**_chat_body(rig, 0), "max_tokens": 64}
return _chat_body(rig, 0)
@pytest.mark.parametrize(
"path",
(
pytest.param("/v1/chat/completions", id="chat"),
pytest.param("/v1/responses", id="responses"),
pytest.param("/v1/messages", id="messages"),
),
)
def test_client_disconnect_mid_tool_call_scans_the_tool_call_it_already_received(
gateway: Gateway, tmp_path: Path, path: str
) -> None:
with _disconnect_rig(gateway, tmp_path, shape="tool_call") as rig:
try:
received: Final = _stream_and_close(rig, path, _tool_call_request(rig, path), rig.token())
assert rig.token() in received, received
scans: Final = eventually(
lambda: rig.scans.response_scans(rig.secret()), lambda values: len(values) >= 1, seconds=4
)
finally:
rig.gate.set()
assert rig.secret() in json.dumps(scans[-1].get("tool_calls")), scans
def test_client_disconnect_after_a_finished_responses_tool_call_scans_the_text_and_tool_call_it_received(
gateway: Gateway, tmp_path: Path
) -> None:
with _disconnect_rig(gateway, tmp_path, shape="text_then_tool_call") as rig:
try:
received: Final = _stream_and_close(
rig, "/v1/responses", _tool_call_request(rig, "/v1/responses"), "response.output_item.done"
)
assert rig.token() in received, received
scans: Final = eventually(
lambda: tuple(scan for scan in rig.scans.response_scans(rig.secret()) if scan.get("texts")),
lambda values: len(values) >= 1,
seconds=4,
)
finally:
rig.gate.set()
assert rig.secret() in json.dumps(scans[-1].get("texts")), scans
assert rig.secret() in json.dumps(scans[-1].get("tool_calls")), scans
def test_client_disconnect_mid_stream_scans_for_a_guardrail_the_request_opted_into(
gateway: Gateway, tmp_path: Path
) -> None:
with _disconnect_rig(gateway, tmp_path, default_on=False) as rig:
body: Final = {**_chat_body(rig, 0), "guardrails": [rig.identity]}
try:
received: Final = _stream_and_close(rig, "/v1/chat/completions", body, rig.secret())
assert rig.secret() in received, received
eventually(lambda: rig.scans.response_scans(rig.secret()), lambda values: len(values) == 1, seconds=4)
finally:
rig.gate.set()
assert _post_call_statuses(rig, rig.model) == (("success",),)
def test_client_disconnect_before_any_content_sends_no_response_scan(gateway: Gateway, tmp_path: Path) -> None:
with _disconnect_rig(gateway, tmp_path, shape="empty") as rig:
try:
_stream_and_close(rig, "/v1/chat/completions", _chat_body(rig, 0), "data: ")
finally:
rig.gate.set()
rows: Final = _post_call_statuses(rig, rig.model)
assert rig.scans.response_scans(rig.secret()) == (), rig.scans.seen
assert len(rows) == 1 and "success" not in rows[0], rows
@pytest.mark.parametrize(
"params",
(
pytest.param(dict(_END_OF_STREAM_ONLY), id="end-of-stream-only"),
pytest.param({"streaming_buffer_until_moderated": True}, id="buffered"),
),
)
def test_a_fully_read_stream_is_scanned_exactly_once(
gateway: Gateway, tmp_path: Path, params: dict[str, JsonValue]
) -> None:
with _disconnect_rig(gateway, tmp_path, params=params, gated=False) as rig:
response: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig, 0))
assert response.status_code == 200, response.text
assert rig.secret() in response.text and "[DONE]" in response.text, response.text
assert _post_call_statuses(rig, rig.model) == (("success",),)
assert len(rig.scans.response_scans(rig.secret())) == 1, rig.scans.seen
def _cached_twin_rows(rig: _DisconnectRig) -> tuple[tuple[str, ...], ...]:
first: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig, 0, stream=False))
second: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig, 0, stream=False))
assert first.status_code == second.status_code == 200, (first.text, second.text)
assert first.json()["choices"] == second.json()["choices"], (first.text, second.text)
assert len(rig.upstream.drain()) == 1, "the second request must be served from the cache"
return _post_call_statuses(rig, rig.model, rows=2)
def test_a_non_streaming_response_and_its_cache_hit_are_each_scanned_once(gateway: Gateway, tmp_path: Path) -> None:
with _disconnect_rig(gateway, tmp_path, gated=False) as rig:
rows: Final = _cached_twin_rows(rig)
assert rows[0] == ("success",), rows
assert len(rig.scans.response_scans(rig.secret())) == 2, (rows, rig.scans.seen)
def test_a_cache_hit_row_records_the_post_call_verdict_of_its_scan(gateway: Gateway, tmp_path: Path) -> None:
pytest.skip("BUG: LIT-9088 the cache-hit spend row drops the post_call verdict of the scan that ran on it")
with _disconnect_rig(gateway, tmp_path, gated=False) as rig:
assert _cached_twin_rows(rig) == (("success",), ("success",))
def test_concurrent_disconnects_during_a_guardrail_outage_each_record_exactly_one_verdict(
gateway: Gateway, tmp_path: Path
) -> None:
def guardrail(request: Request) -> Reply:
assert request.target == "/beta/litellm_basic_guardrail_api", request.target
body: Final = json.loads(request.body)
index: Final = int(_secret_for(request).rsplit("-", 1)[1])
if body["input_type"] == "response" and index % 3 == 0:
return Reply(status=503, body=b'{"error": "synthetic guardrail outage"}')
return Reply(body=json.dumps({"action": "NONE"}).encode())
clients: Final = (_chat_httpx, _responses_httpx, _messages_httpx)
with _disconnect_rig(gateway, tmp_path, guardrail=guardrail, gated=False, pause=3, workers=2) as rig:
with ThreadPoolExecutor(max_workers=30) as pool:
received: Final = tuple(pool.map(lambda index: clients[index % 3](rig, index), range(30)))
assert all(rig.secret(index) in line for index, line in enumerate(received)), received
scanned: Final = eventually(
lambda: tuple(len(rig.scans.response_scans(rig.secret(index) + '"')) for index in range(30)),
lambda counts: all(count >= 1 for count in counts),
seconds=20,
)
assert scanned == (1,) * 30, scanned
chat_rows: Final = _post_call_statuses(rig, rig.model, rows=10)
assert sorted(chat_rows) == sorted(
("guardrail_failed_to_respond",) if index % 3 == 0 else ("success",) for index in range(0, 30, 3)
), chat_rows
def test_disconnect_scans_keep_recording_after_a_worker_is_killed(gateway: Gateway, tmp_path: Path) -> None:
with _disconnect_rig(gateway, tmp_path, gated=False, pause=3, workers=2) as rig:
members: Final = tuple(
member for member in group_members(rig.owned.process.pid) if member.pid != rig.owned.process.pid
)
workers: Final = tuple(member for member in members if any("spawn_main" in part for part in member.cmdline()))
assert len(workers) >= 2, members
workers[0].send_signal(signal.SIGKILL)
psutil.wait_procs((workers[0],), timeout=10)
with ThreadPoolExecutor(max_workers=8) as pool:
received: Final = tuple(pool.map(lambda index: _chat_httpx(rig, index), range(8)))
assert all(rig.secret(index) in line for index, line in enumerate(received)), received
assert _post_call_statuses(rig, rig.model, rows=8) == (("success",),) * 8
assert rig.owned.process.poll() is None

View file

@ -2489,6 +2489,28 @@ class TestAnthropicMessagesHandlerStreamingScanKey:
assert ended_key.tool_calls_in_flight is False
assert ended_key != open_key
def test_released_stream_as_ended_keys_the_tool_use_the_client_already_received(self):
handler = AnthropicMessagesHandler()
tool_use = self._sse(
"content_block_start",
{
"type": "content_block_start",
"index": 1,
"content_block": {"type": "tool_use", "id": "toolu_1", "name": "get_weather", "input": {}},
},
)
stopped_key = handler.get_streaming_scan_key([self._text_delta("hi"), tool_use, self._stop("tool_use")])
released_key = handler.get_streaming_scan_key(
handler.released_stream_as_ended([self._text_delta("hi"), tool_use])
)
assert released_key.stream_ended is True
assert released_key == stopped_key
def test_released_stream_as_ended_leaves_a_text_only_stream_as_released(self):
released = (self._text_delta("hi"), self._text_delta(" there"))
ended = AnthropicMessagesHandler().released_stream_as_ended(released)
assert ended == released and all(a is b for a, b in zip(ended, released, strict=True))
class PerRowTextGuardrail(CustomGuardrail):
"""Answers one redacted text per chat row it was shown, the way a guardrail

View file

@ -2336,3 +2336,32 @@ class TestStreamingScanKey:
handler = OpenAIChatCompletionsHandler()
key = handler.get_streaming_scan_key([self._chunk("hi"), b"data: [DONE]"])
assert key.texts == ("hi",)
def test_released_stream_as_ended_finishes_only_the_choice_whose_tool_call_was_in_flight(self):
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices
tool_call = {"index": 0, "id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}}
released = (
self._chunk("hi", index=0),
ModelResponseStream(choices=[StreamingChoices(index=1, delta=Delta(tool_calls=[tool_call]))]),
)
handler = OpenAIChatCompletionsHandler()
ended = handler.released_stream_as_ended(released)
assert all(a is b for a, b in zip(ended[:-1], released, strict=True))
assert [(choice.index, choice.finish_reason) for choice in ended[-1].choices] == [(1, "tool_calls")]
ended_key = handler.get_streaming_scan_key(ended)
assert ended_key.stream_ended is True and len(ended_key.tool_calls) == 1, ended_key
def test_released_stream_as_ended_leaves_a_stream_with_no_tool_call_in_flight_as_released(self):
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices
tool_call = {"index": 0, "id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}}
text_only = (self._chunk("hi", index=0), self._chunk(" there", index=1))
finished_tool_call = (
ModelResponseStream(choices=[StreamingChoices(index=0, delta=Delta(tool_calls=[tool_call]))]),
self._chunk(None, finish_reason="tool_calls", index=0),
)
handler = OpenAIChatCompletionsHandler()
for released in (text_only, finished_tool_call):
ended = handler.released_stream_as_ended(released)
assert len(ended) == len(released) and all(a is b for a, b in zip(ended, released, strict=True))

View file

@ -6,6 +6,7 @@ with guardrail transformations.
"""
import copy
import json
from collections.abc import Callable
from typing import Any, Final, List, Literal, Optional, Tuple
from unittest.mock import AsyncMock, MagicMock, patch
@ -3654,3 +3655,85 @@ class TestOpenAIResponsesHandlerStreamingScanKey:
ended_key = handler.get_streaming_scan_key([self._delta(0, "hi"), added, self._completed(3, [function_call])])
assert ended_key.tool_calls_in_flight is False
assert len(ended_key.tool_calls) == 1
def test_released_stream_as_ended_keys_the_tool_call_the_client_already_received(self):
handler = OpenAIResponsesHandler()
added = {
"type": "response.output_item.added",
"sequence_number": 1,
"item": {"type": "function_call", "id": "fc_1", "call_id": "call_1", "name": "get_weather", "arguments": ""},
}
arguments_delta = {
"type": "response.function_call_arguments.delta",
"sequence_number": 2,
"item_id": "fc_1",
"delta": '{"city": "Paris"',
}
ended_key = handler.get_streaming_scan_key(
handler.released_stream_as_ended([self._delta(0, "hi"), added, arguments_delta])
)
assert ended_key.stream_ended is True
assert ended_key.texts == ("hi",)
assert len(ended_key.tool_calls) == 1 and "Paris" in ended_key.tool_calls[0], ended_key
def test_released_stream_as_ended_leaves_a_text_only_stream_as_released(self):
released = (self._delta(0, "hi"), self._delta(1, " there"))
ended = OpenAIResponsesHandler().released_stream_as_ended(released)
assert ended == released and all(a is b for a, b in zip(ended, released, strict=True))
@staticmethod
def _finished_function_call(sequence_number: int, item_id: str, city: str) -> tuple[dict[str, object], ...]:
arguments = json.dumps({"city": city})
pending = {"type": "function_call", "id": item_id, "call_id": "call_" + item_id, "name": "get_weather"}
return (
{"type": "response.output_item.added", "sequence_number": sequence_number, "item": {**pending, "arguments": ""}},
{
"type": "response.function_call_arguments.delta",
"sequence_number": sequence_number + 1,
"item_id": item_id,
"delta": arguments,
},
{
"type": "response.output_item.done",
"sequence_number": sequence_number + 2,
"item": {**pending, "arguments": arguments, "status": "completed"},
},
)
def test_released_stream_as_ended_keys_every_tool_call_finished_before_the_disconnect(self):
handler = OpenAIResponsesHandler()
released = (
self._delta(0, "hi"),
*self._finished_function_call(1, "fc_1", "Paris"),
*self._finished_function_call(4, "fc_2", "Rome"),
)
ended_key = handler.get_streaming_scan_key(handler.released_stream_as_ended(released))
assert ended_key.stream_ended is True
assert ended_key.texts == ("hi",)
cities = tuple(city for fingerprint in ended_key.tool_calls for city in ("Paris", "Rome") if city in fingerprint)
assert cities == ("Paris", "Rome"), ended_key
def test_released_stream_as_ended_keys_a_message_whose_item_already_finished(self):
handler = OpenAIResponsesHandler()
message_done = {
"type": "response.output_item.done",
"sequence_number": 1,
"item": {"type": "message", "id": "msg_1", "content": [{"type": "output_text", "text": "hi"}]},
}
ended_key = handler.get_streaming_scan_key(handler.released_stream_as_ended((self._delta(0, "hi"), message_done)))
assert ended_key.stream_ended is True
assert ended_key.texts == ("hi",)
@pytest.mark.asyncio
async def test_scan_of_a_stream_released_through_a_finished_tool_call_covers_its_text_too(self):
handler = OpenAIResponsesHandler()
guardrail = MockRecordingGuardrail(guardrail_name="test")
released = (self._delta(0, "hi"), *self._finished_function_call(1, "fc_1", "Paris"))
await handler.process_output_streaming_response(
responses_so_far=list(handler.released_stream_as_ended(released)),
guardrail_to_apply=guardrail,
request_data={},
)
assert [inputs.get("texts") for inputs in guardrail.seen_inputs] == [["hi"]], guardrail.seen_inputs
tool_calls = guardrail.seen_inputs[0].get("tool_calls") or []
assert [call["function"]["arguments"] for call in tool_calls] == ['{"city": "Paris"}'], tool_calls

View file

@ -5707,6 +5707,41 @@ async def test_unbuffered_end_of_stream_hook_yields_chunks_before_scan():
assert len(chunk_events) == 3
@pytest.mark.asyncio
async def test_unbuffered_end_of_stream_hook_scans_released_chunks_when_the_client_closes_early():
guardrail = BedrockGuardrail(
guardrail_name="bedrock-audit-mode",
guardrailIdentifier="test-id",
guardrailVersion="DRAFT",
event_hook=GuardrailEventHooks.post_call,
default_on=True,
streaming_buffer_until_moderated=False,
streaming_end_of_stream_only=True,
)
scans = []
async def record_scan(*args, **kwargs):
scans.append(kwargs["source"])
return {"action": "NONE", "assessments": [], "outputs": []}
async def mock_stream():
yield _chat_chunk("Hello", None)
yield _chat_chunk(" world", None)
yield _chat_chunk("", "stop")
with patch.object(guardrail, "make_bedrock_api_request", AsyncMock(side_effect=record_scan)):
stream = guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=UserAPIKeyAuth(),
response=mock_stream(),
request_data={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hi"}]},
)
first = await stream.__anext__()
await stream.aclose()
assert first.choices[0].delta.content == "Hello"
assert scans == ["OUTPUT"]
@pytest.mark.asyncio
async def test_buffered_default_hook_scans_before_any_chunk():
guardrail = BedrockGuardrail(

View file

@ -1,16 +1,21 @@
"""Tests for unified guardrail."""
import asyncio
import contextlib
import io
import logging
from collections.abc import AsyncGenerator, AsyncIterable, AsyncIterator
from types import SimpleNamespace
from typing import TYPE_CHECKING, Final, Literal
import anyio
import pytest
import litellm
from litellm.caching import DualCache
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
ModifyResponseException,
log_guardrail_information,
)
from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route
@ -42,7 +47,15 @@ from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrai
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.utils import CallTypes, Delta, GenericGuardrailAPIInputs, ModelResponseStream, StreamingChoices
from litellm.types.utils import (
CallTypes,
Delta,
GenericGuardrailAPIInputs,
ModelResponse,
ModelResponseStream,
StandardLoggingGuardrailInformation,
StreamingChoices,
)
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@ -2164,6 +2177,554 @@ class _ScanCountingGuardrail(CustomGuardrail):
return inputs
class _GatedScanGuardrail(_ScanCountingGuardrail):
"""End-of-stream scan that holds until released, recording scans that finished."""
def __init__(self) -> None:
super().__init__(end_of_stream_only=True)
self.scan_started = anyio.Event()
self.scan_released = anyio.Event()
self.finished_scans = 0
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict[str, object],
input_type: Literal["request", "response"],
**kwargs: object,
) -> GenericGuardrailAPIInputs:
self.scan_started.set()
await self.scan_released.wait()
recorded = await super().apply_guardrail(inputs, request_data, input_type, **kwargs)
self.finished_scans += 1
return recorded
class _FinishReasonRecordingGuardrail(_ScanCountingGuardrail):
"""Records the finish reasons of the stream handed to each response-side scan"""
def __init__(self) -> None:
super().__init__(end_of_stream_only=True)
self.finish_reasons: tuple[str | None, ...] = ()
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict[str, object],
input_type: Literal["request", "response"],
**kwargs: object,
) -> GenericGuardrailAPIInputs:
rebuilt = request_data.get("response")
released = request_data.get("responses")
chunks = (
tuple(chunk for chunk in released if isinstance(chunk, ModelResponseStream))
if isinstance(released, list)
else ()
)
choices = [
*(rebuilt.choices if isinstance(rebuilt, ModelResponse) else ()),
*(choice for chunk in chunks for choice in chunk.choices),
]
self.finish_reasons = (*self.finish_reasons, *(choice.finish_reason for choice in choices))
return await super().apply_guardrail(inputs, request_data, input_type, **kwargs)
class _GatedToolCallGuardrail(_StreamingTextGuardrail):
"""Tool-call inspection that holds until released"""
def __init__(self) -> None:
super().__init__()
self.inspection_started = anyio.Event()
self.inspection_released = anyio.Event()
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict[str, object],
input_type: Literal["request", "response"],
**kwargs: object,
) -> GenericGuardrailAPIInputs:
if input_type == "response" and inputs.get("tool_calls"):
self.inspection_started.set()
await self.inspection_released.wait()
return await super().apply_guardrail(inputs, request_data, input_type, **kwargs)
class _MarkerBlockingScanGuardrail(_ScanCountingGuardrail):
"""Scan-counting guardrail that blocks any scan whose text contains BLOCKME"""
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict[str, object],
input_type: Literal["request", "response"],
**kwargs: object,
) -> GenericGuardrailAPIInputs:
recorded = await super().apply_guardrail(inputs, request_data, input_type, **kwargs)
if any("BLOCKME" in text for text in recorded.get("texts") or []):
raise ModifyResponseException(
message="blocked", model="gpt-4", request_data=request_data, guardrail_name=self.guardrail_name
)
return recorded
class _MarkerHttpErrorScanGuardrail(_ScanCountingGuardrail):
"""Scan-counting guardrail that raises an HTTPException for any scan whose text contains BLOCKME"""
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict[str, object],
input_type: Literal["request", "response"],
**kwargs: object,
) -> GenericGuardrailAPIInputs:
recorded = await super().apply_guardrail(inputs, request_data, input_type, **kwargs)
if any("BLOCKME" in text for text in recorded.get("texts") or []):
raise unified_module.HTTPException(status_code=400, detail={"error": "Violated guardrail policy"})
return recorded
class _MarkerBlockingStreamingTextGuardrail(_StreamingTextGuardrail):
"""incremental_diff guardrail that blocks any round whose text contains BLOCKME"""
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict[str, object],
input_type: Literal["request", "response"],
**kwargs: object,
) -> GenericGuardrailAPIInputs:
transformed = await super().apply_guardrail(inputs, request_data, input_type, **kwargs)
if any("BLOCKME" in text for text in inputs.get("texts") or []):
raise ModifyResponseException(
message="blocked", model="gpt-4", request_data=request_data, guardrail_name=self.guardrail_name
)
return transformed
class _DisconnectRewritingGuardrail(_ScanCountingGuardrail):
"""End-of-stream guardrail that rewrites every scanned text"""
def __init__(self) -> None:
super().__init__(end_of_stream_only=True)
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict[str, object],
input_type: Literal["request", "response"],
**kwargs: object,
) -> GenericGuardrailAPIInputs:
recorded = await super().apply_guardrail(inputs, request_data, input_type, **kwargs)
return {**recorded, "texts": ["REWRITTEN" for _ in recorded.get("texts") or []]}
class _RecordedScanGuardrail(CustomGuardrail):
"""End-of-stream scan recorded through log_guardrail_information, returning ``reply`` or raising ``error``"""
def __init__(self, *, reply: GenericGuardrailAPIInputs | None = None, error: Exception | None = None) -> None:
super().__init__(guardrail_name="recorded-scan")
self.streaming_end_of_stream_only = True
self.streaming_buffer_until_moderated = False
self.guardrail_config = {}
self._reply = reply
self._error = error
def should_run_guardrail(self, data: dict[str, object], event_type: GuardrailEventHooks) -> bool:
return True
@log_guardrail_information
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict[str, object],
input_type: Literal["request", "response"],
**kwargs: object,
) -> GenericGuardrailAPIInputs:
if self._error is not None:
raise self._error
return inputs if self._reply is None else self._reply
def _recorded_guardrail_statuses(request_data: dict[str, object]) -> list[str]:
metadata = request_data["metadata"]
assert isinstance(metadata, dict), request_data
entries: list[StandardLoggingGuardrailInformation] = metadata.get("standard_logging_guardrail_information", [])
return [entry["guardrail_status"] for entry in entries]
class TestStreamingClientDisconnectScan:
"""A client that reads streamed content and then disconnects must not skip
the end-of-stream scan of what it already received."""
@pytest.fixture(autouse=True)
def _use_real_mappings(self, monkeypatch: pytest.MonkeyPatch) -> None:
_patch_translation_mappings(monkeypatch, load_guardrail_translation_mappings())
@staticmethod
def _guarded_stream(guardrail: CustomGuardrail, upstream: AsyncIterable[object]) -> AsyncGenerator[object, None]:
return UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test-key", request_route="/v1/chat/completions"),
response=upstream,
request_data={"guardrail_to_apply": guardrail, "model": "gpt-4", "metadata": {}},
)
@pytest.mark.asyncio
async def test_closing_after_released_content_still_scans_it(self):
guardrail = _ScanCountingGuardrail(end_of_stream_only=True)
async def upstream() -> AsyncIterator[ModelResponseStream]:
yield _stream_chunk("synthetic secret")
yield _stream_chunk(" tail", finish_reason="stop")
stream = self._guarded_stream(guardrail, upstream())
received = await stream.__anext__()
await stream.aclose()
assert _delta_text(received) == "synthetic secret"
assert [scan["texts"] for scan in guardrail.scans] == [["synthetic secret"]], guardrail.scans
@pytest.mark.asyncio
async def test_closing_mid_text_stream_does_not_hand_the_scan_a_tool_calls_finish(self):
guardrail = _FinishReasonRecordingGuardrail()
async def upstream() -> AsyncIterator[ModelResponseStream]:
yield _stream_chunk("synthetic secret")
yield _stream_chunk(" tail", finish_reason="stop")
stream = self._guarded_stream(guardrail, upstream())
await stream.__anext__()
await stream.aclose()
assert [scan["texts"] for scan in guardrail.scans] == [["synthetic secret"]], guardrail.scans
assert "tool_calls" not in guardrail.finish_reasons, guardrail.finish_reasons
@pytest.mark.asyncio
async def test_upstream_cancellation_after_released_content_still_scans_it(self):
guardrail = _ScanCountingGuardrail(end_of_stream_only=True)
async def upstream() -> AsyncIterator[ModelResponseStream]:
yield _stream_chunk("synthetic secret")
raise asyncio.CancelledError()
stream = self._guarded_stream(guardrail, upstream())
received = await stream.__anext__()
with pytest.raises(asyncio.CancelledError):
await stream.__anext__()
assert _delta_text(received) == "synthetic secret"
assert [scan["texts"] for scan in guardrail.scans] == [["synthetic secret"]], guardrail.scans
@pytest.mark.asyncio
async def test_cancellation_while_the_disconnect_scan_is_in_flight_lets_it_finish(self):
guardrail = _GatedScanGuardrail()
first_chunk_received = anyio.Event()
async def upstream() -> AsyncIterator[ModelResponseStream]:
yield _stream_chunk("synthetic secret")
await anyio.sleep_forever()
yield _stream_chunk(" tail", finish_reason="stop")
async def consume(scope_ready: list[anyio.CancelScope]) -> None:
with anyio.CancelScope() as scope:
scope_ready.append(scope)
async with contextlib.aclosing(self._guarded_stream(guardrail, upstream())) as stream:
async for _item in stream:
first_chunk_received.set()
scopes = []
with anyio.fail_after(5):
async with anyio.create_task_group() as task_group:
task_group.start_soon(consume, scopes)
await first_chunk_received.wait()
scopes[0].cancel()
await guardrail.scan_started.wait()
await anyio.sleep(0)
guardrail.scan_released.set()
assert guardrail.finished_scans == 1
assert [scan["texts"] for scan in guardrail.scans] == [["synthetic secret"]], guardrail.scans
@pytest.mark.asyncio
async def test_closing_after_a_sampled_scan_covered_everything_released_does_not_scan_again(self):
guardrail = _ScanCountingGuardrail(sampling_rate=1)
async def upstream() -> AsyncIterator[ModelResponseStream]:
yield _stream_chunk("synthetic")
yield _stream_chunk(" secret")
await anyio.sleep_forever()
stream = self._guarded_stream(guardrail, upstream())
released = [await stream.__anext__(), await stream.__anext__()]
await stream.aclose()
scanned_texts = [scan["texts"] for scan in guardrail.scans]
assert "".join(_delta_text(chunk) for chunk in released) == "synthetic secret"
assert scanned_texts[-1] == ["synthetic secret"], scanned_texts
assert len(scanned_texts) == len({tuple(texts) for texts in scanned_texts}), scanned_texts
@pytest.mark.asyncio
async def test_closing_before_any_content_is_released_does_not_scan(self):
guardrail = _ScanCountingGuardrail(end_of_stream_only=True, buffer_until_moderated=True)
upstream_started = anyio.Event()
async def upstream() -> AsyncIterator[ModelResponseStream]:
yield _stream_chunk("withheld")
upstream_started.set()
await anyio.sleep_forever()
yield _stream_chunk("never", finish_reason="stop")
stream = self._guarded_stream(guardrail, upstream())
async with anyio.create_task_group() as task_group:
task_group.start_soon(stream.__anext__)
await upstream_started.wait()
task_group.cancel_scope.cancel()
assert guardrail.scans == ()
@pytest.mark.asyncio
async def test_cancellation_during_end_of_stream_scan_lets_the_scan_finish(self):
guardrail = _GatedScanGuardrail()
async def upstream() -> AsyncIterator[ModelResponseStream]:
yield _stream_chunk("synthetic secret")
yield _stream_chunk(" tail", finish_reason="stop")
async def consume(scope_ready: list[anyio.CancelScope]) -> None:
with anyio.CancelScope() as scope:
scope_ready.append(scope)
async for _item in self._guarded_stream(guardrail, upstream()):
pass
scopes = []
async with anyio.create_task_group() as task_group:
task_group.start_soon(consume, scopes)
await guardrail.scan_started.wait()
scopes[0].cancel()
await anyio.sleep(0)
guardrail.scan_released.set()
assert guardrail.finished_scans == 1
assert [scan["texts"] for scan in guardrail.scans] == [["synthetic secret tail"]], guardrail.scans
@staticmethod
async def _close_after_first_chunk(guardrail: CustomGuardrail) -> dict[str, object]:
request_data: dict[str, object] = {"guardrail_to_apply": guardrail, "model": "gpt-4", "metadata": {}}
async def upstream() -> AsyncIterator[ModelResponseStream]:
yield _stream_chunk("synthetic secret")
yield _stream_chunk(" tail", finish_reason="stop")
stream = UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test-key", request_route="/v1/chat/completions"),
response=upstream(),
request_data=request_data,
)
received = await stream.__anext__()
await stream.aclose()
assert _delta_text(received) == "synthetic secret"
return request_data
@pytest.mark.asyncio
async def test_disconnect_scan_that_fails_after_the_verdict_records_the_failure(self):
request_data = await self._close_after_first_chunk(
_RecordedScanGuardrail(reply={"texts": ["synthetic secret", "unmatched extra text"]})
)
assert _recorded_guardrail_statuses(request_data) == ["success", "guardrail_failed_to_respond"]
@pytest.mark.asyncio
async def test_disconnect_scan_whose_guardrail_raises_records_one_failure(self):
request_data = await self._close_after_first_chunk(_RecordedScanGuardrail(error=RuntimeError("provider down")))
assert _recorded_guardrail_statuses(request_data) == ["guardrail_failed_to_respond"]
@pytest.mark.asyncio
async def test_disconnect_scan_that_passes_records_only_the_verdict(self):
request_data = await self._close_after_first_chunk(_RecordedScanGuardrail())
assert _recorded_guardrail_statuses(request_data) == ["success"]
@staticmethod
async def _tool_call_upstream() -> AsyncIterator[ModelResponseStream]:
from litellm.types.utils import ChatCompletionDeltaToolCall, Function
tool_call = ChatCompletionDeltaToolCall(
id="call_1", index=0, type="function", function=Function(name="get_weather", arguments='{"city": "Paris"}')
)
yield ModelResponseStream(
choices=[StreamingChoices(index=0, delta=Delta(content=None, tool_calls=[tool_call]))]
)
yield _stream_chunk(None, finish_reason="tool_calls")
@pytest.mark.asyncio
async def test_cancellation_during_incremental_diff_tool_call_inspection_lets_it_finish_once(self):
guardrail = _GatedToolCallGuardrail()
async def consume(scope_ready: list[anyio.CancelScope]) -> None:
with anyio.CancelScope() as scope:
scope_ready.append(scope)
async with contextlib.aclosing(self._guarded_stream(guardrail, self._tool_call_upstream())) as stream:
async for _item in stream:
pass
scopes = []
async with anyio.create_task_group() as task_group:
task_group.start_soon(consume, scopes)
await guardrail.inspection_started.wait()
scopes[0].cancel()
await anyio.sleep(0)
guardrail.inspection_released.set()
assert [[call["function"]["name"] for call in calls] for calls in guardrail.received_tool_calls] == [
["get_weather"]
], guardrail.received_tool_calls
@pytest.mark.asyncio
async def test_closing_after_the_incremental_diff_tool_call_inspection_does_not_inspect_again(self):
guardrail = _StreamingTextGuardrail()
stream = self._guarded_stream(guardrail, self._tool_call_upstream())
released = [await stream.__anext__(), await stream.__anext__()]
await stream.aclose()
assert released[-1].choices[0].finish_reason == "tool_calls"
assert [[call["function"]["name"] for call in calls] for calls in guardrail.received_tool_calls] == [
["get_weather"]
], guardrail.received_tool_calls
@pytest.mark.asyncio
async def test_closing_after_a_released_tool_call_under_incremental_diff_still_inspects_it(self):
guardrail = _StreamingTextGuardrail()
stream = self._guarded_stream(guardrail, self._tool_call_upstream())
received = await stream.__anext__()
await stream.aclose()
assert [call.function.name for call in received.choices[0].delta.tool_calls] == ["get_weather"]
assert [[call["function"]["name"] for call in calls] for calls in guardrail.received_tool_calls] == [
["get_weather"]
], guardrail.received_tool_calls
@pytest.mark.asyncio
async def test_closing_while_a_mid_stream_block_is_delivered_does_not_scan_the_blocked_content_again(self):
guardrail = _MarkerBlockingScanGuardrail(sampling_rate=2)
async def upstream() -> AsyncIterator[ModelResponseStream]:
for text in ("a", "b", "c", "BLOCKME"):
yield _stream_chunk(text)
yield _stream_chunk(" tail", finish_reason="stop")
stream = self._guarded_stream(guardrail, upstream())
received = [_delta_text(await stream.__anext__()) for _ in range(3)]
await stream.__anext__()
await stream.aclose()
assert received == ["a", "b", "c"]
assert [scan["texts"] for scan in guardrail.scans] == [["ab"], ["abcBLOCKME"]], guardrail.scans
@pytest.mark.asyncio
async def test_closing_while_a_mid_stream_guardrail_error_is_delivered_does_not_scan_the_content_again(self):
guardrail = _MarkerHttpErrorScanGuardrail(sampling_rate=2)
async def upstream() -> AsyncIterator[ModelResponseStream]:
for text in ("a", "b", "c", "BLOCKME"):
yield _stream_chunk(text)
yield _stream_chunk(" tail", finish_reason="stop")
stream = self._guarded_stream(guardrail, upstream())
received = [_delta_text(await stream.__anext__()) for _ in range(3)]
await stream.__anext__()
await stream.aclose()
assert received == ["a", "b", "c"]
assert [scan["texts"] for scan in guardrail.scans] == [["ab"], ["abcBLOCKME"]], guardrail.scans
@pytest.mark.asyncio
async def test_closing_while_the_scanned_final_chunk_is_delivered_does_not_scan_the_stream_again(self):
guardrail = _ScanCountingGuardrail(end_of_stream_only=True)
async def upstream() -> AsyncIterator[ModelResponseStream]:
yield _stream_chunk("a")
yield _stream_chunk("b")
yield _stream_chunk(" tail", finish_reason="stop")
stream = self._guarded_stream(guardrail, upstream())
received = [_delta_text(await stream.__anext__()) for _ in range(3)]
await stream.aclose()
assert received == ["a", "b", " tail"]
assert [scan["texts"] for scan in guardrail.scans] == [["ab tail"]], guardrail.scans
@pytest.mark.asyncio
async def test_cancellation_with_a_withheld_window_scans_only_the_released_chunks(self):
guardrail = _ScanCountingGuardrail(sampling_rate=2, buffer_until_moderated=True)
guardrail.streaming_buffer_release_on_scan = True
async def upstream() -> AsyncIterator[ModelResponseStream]:
yield _stream_chunk("a")
yield _stream_chunk("b")
yield _stream_chunk("WITHHELD")
raise asyncio.CancelledError()
stream = self._guarded_stream(guardrail, upstream())
received = [_delta_text(await stream.__anext__()), _delta_text(await stream.__anext__())]
with pytest.raises(asyncio.CancelledError):
await stream.__anext__()
assert received == ["a", "b"]
assert [scan["texts"] for scan in guardrail.scans] == [["ab"]], guardrail.scans
@pytest.mark.asyncio
async def test_disconnect_scan_that_rewrites_text_leaves_the_released_chunks_untouched(self):
stream = self._guarded_stream(_DisconnectRewritingGuardrail(), self._text_then_tail_upstream())
received = await stream.__anext__()
await stream.aclose()
assert _delta_text(received) == "synthetic secret"
@staticmethod
async def _text_then_tail_upstream() -> AsyncIterator[ModelResponseStream]:
yield _stream_chunk("synthetic secret")
yield _stream_chunk(" tail", finish_reason="stop")
@pytest.mark.asyncio
async def test_closing_while_an_incremental_diff_block_is_delivered_does_not_inspect_the_tool_call_again(self):
guardrail = _MarkerBlockingStreamingTextGuardrail()
async def upstream() -> AsyncIterator[ModelResponseStream]:
async for tool_chunk in self._tool_call_upstream():
if tool_chunk.choices[0].finish_reason is None:
yield tool_chunk
yield _stream_chunk("BLOCKME")
yield _stream_chunk(" tail", finish_reason="stop")
stream = self._guarded_stream(guardrail, upstream())
released = [await stream.__anext__(), await stream.__anext__()]
await stream.aclose()
assert [call.function.name for call in released[0].choices[0].delta.tool_calls] == ["get_weather"]
assert guardrail.received_texts == [["BLOCKME"]], guardrail.received_texts
assert guardrail.received_tool_calls == [], guardrail.received_tool_calls
@pytest.mark.asyncio
async def test_closing_after_a_released_tool_call_under_incremental_diff_scans_no_held_back_text(self):
guardrail = _StreamingTextGuardrail(holdback_schedule=[len("held secret")] * 2)
async def upstream() -> AsyncIterator[ModelResponseStream]:
yield _stream_chunk("held secret")
async for tool_chunk in self._tool_call_upstream():
yield tool_chunk
stream = self._guarded_stream(guardrail, upstream())
received = await stream.__anext__()
await stream.aclose()
assert [call.function.name for call in received.choices[0].delta.tool_calls] == ["get_weather"]
assert guardrail.received_tool_calls, guardrail.received_texts
assert all("held secret" not in text for text in guardrail.received_texts[-1]), guardrail.received_texts
def _responses_delta(sequence_number, text):
return {
"type": "response.output_text.delta",

View file

@ -7,6 +7,7 @@ Verifies that the hook:
3. Actually yields chunks from async generators
"""
import logging
from typing import AsyncGenerator, Any
from unittest.mock import MagicMock, patch
@ -31,7 +32,7 @@ class MockStreamingCallback(CustomLogger):
self,
user_api_key_dict: UserAPIKeyAuth,
response: AsyncGenerator[Any, None],
request_data: dict,
request_data: dict[str, object],
) -> AsyncGenerator[Any, None]:
"""Transform chunks by tracking and optionally prefixing."""
async for chunk in response:
@ -185,3 +186,340 @@ async def test_streaming_hook_propagates_callback_errors():
with pytest.raises(RuntimeError, match="Callback failed!"):
async for _ in result:
pass
class CleanupRecordingCallback(CustomLogger):
"""Iterator hook whose cleanup marks when it ran."""
def __init__(self):
super().__init__()
self.cleaned_up = False
async def async_post_call_streaming_iterator_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
response: AsyncGenerator[Any, None],
request_data: dict[str, object],
) -> AsyncGenerator[Any, None]:
try:
async for chunk in response:
yield chunk
finally:
self.cleaned_up = True
@pytest.mark.asyncio
async def test_closing_the_stream_runs_every_callback_cleanup_before_returning():
proxy_logging = ProxyLogging(user_api_key_cache=MagicMock())
callbacks = [CleanupRecordingCallback(), CleanupRecordingCallback()]
with patch.object(litellm, "callbacks", callbacks):
ProxyLogging._callback_capabilities_cache.clear()
stream = proxy_logging.async_post_call_streaming_iterator_hook(
response=mock_streaming_response(),
user_api_key_dict=UserAPIKeyAuth(api_key="test_key"),
request_data={"model": "gpt-4", "messages": []},
)
first = await stream.__anext__()
await stream.aclose()
ProxyLogging._callback_capabilities_cache.clear()
assert first == {"choices": [{"delta": {"content": "Hello"}}]}
assert [callback.cleaned_up for callback in callbacks] == [True, True]
class RaisingCleanupCallback(CustomLogger):
"""Iterator hook whose cleanup raises."""
async def async_post_call_streaming_iterator_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
response: AsyncGenerator[Any, None],
request_data: dict[str, object],
) -> AsyncGenerator[Any, None]:
try:
async for chunk in response:
yield chunk
finally:
raise RuntimeError("cleanup failed")
@pytest.mark.asyncio
async def test_closing_the_stream_still_cleans_up_inner_callbacks_when_an_outer_cleanup_raises():
proxy_logging = ProxyLogging(user_api_key_cache=MagicMock())
inner = CleanupRecordingCallback()
with patch.object(litellm, "callbacks", [inner, RaisingCleanupCallback()]):
ProxyLogging._callback_capabilities_cache.clear()
stream = proxy_logging.async_post_call_streaming_iterator_hook(
response=mock_streaming_response(),
user_api_key_dict=UserAPIKeyAuth(api_key="test_key"),
request_data={"model": "gpt-4", "messages": []},
)
await stream.__anext__()
await stream.aclose()
ProxyLogging._callback_capabilities_cache.clear()
assert inner.cleaned_up is True
class _PlainAsyncIterator:
def __init__(self, response: AsyncGenerator[Any, None]) -> None:
self._response = response
def __aiter__(self) -> "_PlainAsyncIterator":
return self
async def __anext__(self) -> Any:
return await self._response.__anext__()
class PlainIteratorCallback(CustomLogger):
"""Iterator hook that returns an async iterator with no aclose."""
def async_post_call_streaming_iterator_hook( # pyright: ignore[reportIncompatibleMethodOverride] # a plain async iterator worked before aclose handling
self,
user_api_key_dict: UserAPIKeyAuth,
response: AsyncGenerator[Any, None],
request_data: dict[str, object],
) -> _PlainAsyncIterator:
return _PlainAsyncIterator(response)
@pytest.mark.asyncio
async def test_a_hook_returning_a_plain_async_iterator_streams_every_chunk():
proxy_logging = ProxyLogging(user_api_key_cache=MagicMock())
with patch.object(litellm, "callbacks", [PlainIteratorCallback()]):
ProxyLogging._callback_capabilities_cache.clear()
received = [
chunk
async for chunk in proxy_logging.async_post_call_streaming_iterator_hook(
response=mock_streaming_response(),
user_api_key_dict=UserAPIKeyAuth(api_key="test_key"),
request_data={"model": "gpt-4", "messages": []},
)
]
ProxyLogging._callback_capabilities_cache.clear()
assert [chunk async for chunk in mock_streaming_response()] == received
class _ClosableAsyncIterator(_PlainAsyncIterator):
def __init__(self, response: AsyncGenerator[Any, None]) -> None:
super().__init__(response)
self.closed = False
async def aclose(self) -> None:
self.closed = True
class _SyncClosableAsyncIterator(_PlainAsyncIterator):
def __init__(self, response: AsyncGenerator[Any, None]) -> None:
super().__init__(response)
self.closed = False
def aclose(self) -> None:
self.closed = True
class ClosableIteratorCallback(CustomLogger):
"""Iterator hook that returns a non-generator async iterator with its own aclose."""
def __init__(
self,
iterator_type: type[_ClosableAsyncIterator] | type[_SyncClosableAsyncIterator] = _ClosableAsyncIterator,
) -> None:
super().__init__()
self.iterator_type = iterator_type
self.returned: tuple[_ClosableAsyncIterator | _SyncClosableAsyncIterator, ...] = ()
def async_post_call_streaming_iterator_hook( # pyright: ignore[reportIncompatibleMethodOverride] # a custom async iterator worked before aclose handling
self,
user_api_key_dict: UserAPIKeyAuth,
response: AsyncGenerator[Any, None],
request_data: dict[str, object],
) -> _ClosableAsyncIterator | _SyncClosableAsyncIterator:
iterator = self.iterator_type(response)
self.returned = (*self.returned, iterator)
return iterator
@pytest.mark.asyncio
async def test_closing_the_stream_closes_a_hook_iterator_that_is_not_a_generator():
proxy_logging = ProxyLogging(user_api_key_cache=MagicMock())
callback = ClosableIteratorCallback()
with patch.object(litellm, "callbacks", [callback]):
ProxyLogging._callback_capabilities_cache.clear()
stream = proxy_logging.async_post_call_streaming_iterator_hook(
response=mock_streaming_response(),
user_api_key_dict=UserAPIKeyAuth(api_key="test_key"),
request_data={"model": "gpt-4", "messages": []},
)
await stream.__anext__()
await stream.aclose()
ProxyLogging._callback_capabilities_cache.clear()
assert [iterator.closed for iterator in callback.returned] == [True]
@pytest.mark.asyncio
async def test_a_hook_iterator_with_a_synchronous_aclose_streams_everything_and_is_closed():
proxy_logging = ProxyLogging(user_api_key_cache=MagicMock())
callback = ClosableIteratorCallback(iterator_type=_SyncClosableAsyncIterator)
with patch.object(litellm, "callbacks", [callback]):
ProxyLogging._callback_capabilities_cache.clear()
received = [
chunk
async for chunk in proxy_logging.async_post_call_streaming_iterator_hook(
response=mock_streaming_response(),
user_api_key_dict=UserAPIKeyAuth(api_key="test_key"),
request_data={"model": "gpt-4", "messages": []},
)
]
ProxyLogging._callback_capabilities_cache.clear()
assert [chunk async for chunk in mock_streaming_response()] == received
assert [iterator.closed for iterator in callback.returned] == [True]
class _RaisingAcloseIterator(_ClosableAsyncIterator):
"""Non-generator async iterator whose asynchronous aclose raises."""
def __init__(self, response: AsyncGenerator[Any, None], error: Exception) -> None:
super().__init__(response)
self.error = error
async def aclose(self) -> None:
self.closed = True
raise self.error
class _SyncRaisingAcloseIterator(_SyncClosableAsyncIterator):
"""Non-generator async iterator whose synchronous aclose raises."""
def __init__(self, response: AsyncGenerator[Any, None], error: Exception) -> None:
super().__init__(response)
self.error = error
def aclose(self) -> None:
self.closed = True
raise self.error
class RaisingAcloseIteratorCallback(CustomLogger):
"""Iterator hook that returns a non-generator async iterator whose aclose raises."""
def __init__(
self,
iterator_type: type[_RaisingAcloseIterator] | type[_SyncRaisingAcloseIterator] = _RaisingAcloseIterator,
error: Exception | None = None,
) -> None:
super().__init__()
self.iterator_type = iterator_type
self.error = error if error is not None else RuntimeError("cleanup failed")
self.returned: tuple[_RaisingAcloseIterator | _SyncRaisingAcloseIterator, ...] = ()
def async_post_call_streaming_iterator_hook( # pyright: ignore[reportIncompatibleMethodOverride] # a custom async iterator worked before aclose handling
self,
user_api_key_dict: UserAPIKeyAuth,
response: AsyncGenerator[Any, None],
request_data: dict[str, object],
) -> _RaisingAcloseIterator | _SyncRaisingAcloseIterator:
iterator = self.iterator_type(response, self.error)
self.returned = (*self.returned, iterator)
return iterator
@pytest.mark.asyncio
async def test_a_hook_iterator_whose_aclose_raises_still_finishes_the_stream(caplog: pytest.LogCaptureFixture) -> None:
proxy_logging = ProxyLogging(user_api_key_cache=MagicMock())
callback = RaisingAcloseIteratorCallback()
with patch.object(litellm, "callbacks", [callback]):
ProxyLogging._callback_capabilities_cache.clear()
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
received = [
chunk
async for chunk in proxy_logging.async_post_call_streaming_iterator_hook(
response=mock_streaming_response(),
user_api_key_dict=UserAPIKeyAuth(api_key="test_key"),
request_data={"model": "gpt-4", "messages": []},
)
]
ProxyLogging._callback_capabilities_cache.clear()
assert [chunk async for chunk in mock_streaming_response()] == received
assert [iterator.closed for iterator in callback.returned] == [True]
warnings_emitted = [
record.getMessage()
for record in caplog.records
if record.levelname == "WARNING" and "RaisingAcloseIteratorCallback" in record.getMessage()
]
assert len(warnings_emitted) == 1
assert "RuntimeError" in warnings_emitted[0]
assert "cleanup failed" not in warnings_emitted[0]
@pytest.mark.asyncio
async def test_a_hook_iterator_whose_synchronous_aclose_raises_still_finishes_the_stream(
caplog: pytest.LogCaptureFixture,
) -> None:
proxy_logging = ProxyLogging(user_api_key_cache=MagicMock())
callback = RaisingAcloseIteratorCallback(
iterator_type=_SyncRaisingAcloseIterator, error=ValueError("sync cleanup failed")
)
with patch.object(litellm, "callbacks", [callback]):
ProxyLogging._callback_capabilities_cache.clear()
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
received = [
chunk
async for chunk in proxy_logging.async_post_call_streaming_iterator_hook(
response=mock_streaming_response(),
user_api_key_dict=UserAPIKeyAuth(api_key="test_key"),
request_data={"model": "gpt-4", "messages": []},
)
]
ProxyLogging._callback_capabilities_cache.clear()
assert [chunk async for chunk in mock_streaming_response()] == received
assert [iterator.closed for iterator in callback.returned] == [True]
warnings_emitted = [
record.getMessage()
for record in caplog.records
if record.levelname == "WARNING" and "RaisingAcloseIteratorCallback" in record.getMessage()
]
assert len(warnings_emitted) == 1
assert "ValueError" in warnings_emitted[0]
@pytest.mark.asyncio
async def test_a_hook_iterator_with_a_clean_aclose_streams_everything_without_warning(
caplog: pytest.LogCaptureFixture,
) -> None:
proxy_logging = ProxyLogging(user_api_key_cache=MagicMock())
callback = ClosableIteratorCallback()
with patch.object(litellm, "callbacks", [callback]):
ProxyLogging._callback_capabilities_cache.clear()
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
received = [
chunk
async for chunk in proxy_logging.async_post_call_streaming_iterator_hook(
response=mock_streaming_response(),
user_api_key_dict=UserAPIKeyAuth(api_key="test_key"),
request_data={"model": "gpt-4", "messages": []},
)
]
ProxyLogging._callback_capabilities_cache.clear()
assert [chunk async for chunk in mock_streaming_response()] == received
assert [iterator.closed for iterator in callback.returned] == [True]
assert not [
record.getMessage()
for record in caplog.records
if record.levelname == "WARNING" and "ClosableIteratorCallback" in record.getMessage()
]

View file

@ -26,6 +26,7 @@ from litellm.constants import (
CLIENT_REQUESTED_MODEL_SCOPE_KEY,
MAX_LITELLM_CALL_ID_LENGTH,
RETURN_RAW_MODEL_NAME_METADATA_KEY,
STREAM_SSE_KEEPALIVE_PING_BYTES,
)
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.opentelemetry import UserAPIKeyAuth
@ -4256,6 +4257,72 @@ class TestStreamCloseOnDisconnect:
assert upstream.aclosed
async def test_async_streaming_data_generator_closes_the_guardrail_chain_on_client_disconnect(
self,
):
cleanup_ran = []
async def guarded_chain(**_kwargs):
try:
yield {"type": "chunk"}
yield {"type": "chunk"}
finally:
cleanup_ran.append(True)
proxy_logging_obj = ProxyLogging(user_api_key_cache=MagicMock())
proxy_logging_obj.async_post_call_streaming_iterator_hook = guarded_chain
gen = ProxyBaseLLMRequestProcessing.async_streaming_data_generator(
response=MagicMock(),
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth),
request_data={"model": "mock-model"},
proxy_logging_obj=proxy_logging_obj,
serialize_chunk=lambda c: "data: x\n\n",
serialize_error=lambda e: "data: error\n\n",
)
await gen.__anext__()
await gen.aclose()
assert cleanup_ran == [True]
async def test_async_streaming_data_generator_refunds_the_budget_when_closing_the_guardrail_chain_raises(
self,
):
async def guarded_chain(**_kwargs):
try:
yield STREAM_SSE_KEEPALIVE_PING_BYTES
yield {"type": "chunk"}
finally:
raise RuntimeError("cleanup failed")
proxy_logging_obj = ProxyLogging(user_api_key_cache=MagicMock())
proxy_logging_obj.async_post_call_streaming_iterator_hook = guarded_chain
reservation = object()
user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
user_api_key_dict.budget_reservation = reservation
gen = ProxyBaseLLMRequestProcessing.async_streaming_data_generator(
response=MagicMock(),
user_api_key_dict=user_api_key_dict,
request_data={"model": "mock-model"},
proxy_logging_obj=proxy_logging_obj,
serialize_chunk=lambda c: "data: x\n\n",
serialize_error=lambda e: "data: error\n\n",
)
released: list[object] = []
async def record_release(budget_reservation: object) -> None:
released.append(budget_reservation)
with patch(
"litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation_on_cancel",
new=record_release,
):
await gen.__anext__()
await gen.aclose()
assert released == [reservation]
async def test_async_streaming_data_generator_redacts_internal_details_on_error(
self,
):

View file

@ -7425,6 +7425,69 @@ async def test_async_data_generator_cleanup_on_early_exit():
mock_response.aclose.assert_awaited_once()
def _guarded_chain_logging(chain):
from litellm.proxy.utils import ProxyLogging
proxy_logging = MagicMock(spec=ProxyLogging)
proxy_logging.async_post_call_streaming_iterator_hook = chain
proxy_logging.async_post_call_streaming_hook = AsyncMock(side_effect=lambda **kwargs: kwargs.get("response"))
proxy_logging.post_call_failure_hook = AsyncMock()
return proxy_logging
@pytest.mark.asyncio
async def test_async_data_generator_closes_the_guardrail_chain_before_returning_on_client_disconnect():
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.proxy_server import async_data_generator
cleanup_ran = []
async def guarded_chain(**_kwargs):
try:
yield {"choices": [{"delta": {"content": "Hello"}}]}
yield {"choices": [{"delta": {"content": " world"}}]}
finally:
cleanup_ran.append(True)
with patch("litellm.proxy.proxy_server.proxy_logging_obj", _guarded_chain_logging(guarded_chain)):
gen = async_data_generator(MagicMock(), MagicMock(spec=UserAPIKeyAuth), {"model": "gpt-4o-mini"})
first_chunk = await gen.__anext__()
await gen.aclose()
assert first_chunk.startswith("data: ")
assert cleanup_ran == [True]
@pytest.mark.asyncio
async def test_async_data_generator_closes_the_guardrail_chain_while_a_keepalive_read_is_pending():
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.proxy_server import async_data_generator
cleanup_ran = []
never_arrives = asyncio.Event()
async def guarded_chain(**_kwargs):
try:
yield {"choices": [{"delta": {"content": "Hello"}}]}
await never_arrives.wait()
yield {"choices": [{"delta": {"content": " world"}}]}
finally:
cleanup_ran.append(True)
with (
patch.object(litellm, "sse_keepalive_ping_interval_seconds", 1.0),
patch("litellm.proxy.proxy_server.proxy_logging_obj", _guarded_chain_logging(guarded_chain)),
):
gen = async_data_generator(MagicMock(), MagicMock(spec=UserAPIKeyAuth), {"model": "gpt-4o-mini"})
first_chunk = await gen.__anext__()
heartbeat = await gen.__anext__()
await gen.aclose()
assert first_chunk.startswith("data: ")
assert heartbeat == ": ping\n\n"
assert cleanup_ran == [True]
@pytest.mark.asyncio
async def test_async_data_generator_uses_direct_stream_fast_path_without_callbacks():
"""

View file

@ -11,9 +11,9 @@ from __future__ import annotations
import asyncio
import json
from collections.abc import AsyncGenerator, Iterator
from copy import deepcopy
import logging
from collections.abc import Iterator
from typing import Any, Callable, Dict, List
from unittest.mock import AsyncMock, MagicMock, patch
@ -26,6 +26,7 @@ from litellm.integrations.custom_guardrail import (
CustomGuardrail,
ModifyResponseException,
)
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.prometheus import PrometheusLogger
from litellm.llms.base_llm.guardrail_translation.utils import stream_item_field
from litellm.proxy._types import UserAPIKeyAuth
@ -2933,3 +2934,62 @@ async def test_streaming_iterator_hook_pipeline_discards_dropped_tool_call_on_re
assert delivered == _responses_function_call_events()
assert any("'gr-post'" in message and "discarded" in message for message in _warnings(caplog))
class _RaisingAcloseIterator:
"""Non-generator async iterator whose aclose raises after the stream ended."""
def __init__(self, response: AsyncGenerator[Any, None]) -> None:
self._response = response
def __aiter__(self) -> "_RaisingAcloseIterator":
return self
async def __anext__(self) -> Any:
return await self._response.__anext__()
async def aclose(self) -> None:
raise RuntimeError("cleanup failed")
class RaisingAcloseCallback(CustomLogger):
"""Iterator hook returning a non-generator async iterator whose aclose raises."""
def async_post_call_streaming_iterator_hook( # pyright: ignore[reportIncompatibleMethodOverride] # a custom async iterator worked before aclose handling
self,
user_api_key_dict: UserAPIKeyAuth,
response: AsyncGenerator[Any, None],
request_data: dict[str, object],
) -> _RaisingAcloseIterator:
return _RaisingAcloseIterator(response)
@pytest.mark.asyncio
async def test_streaming_iterator_hook_pipeline_releases_buffered_content_when_a_callback_aclose_raises(
proxy_logging: ProxyLogging,
make_user_api_key_auth: Callable[..., UserAPIKeyAuth],
monkeypatch: pytest.MonkeyPatch,
caplog: pytest.LogCaptureFixture,
) -> None:
monkeypatch.setattr(
litellm, "callbacks", [_rewriting_stream_guardrail(lambda inputs: {}), RaisingAcloseCallback()]
)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False)
data = _post_call_pipeline_data(stream=True)
chunks = _tool_call_stream_chunks()
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
delivered = [
item
async for item in proxy_logging.async_post_call_streaming_iterator_hook(
user_api_key_dict=make_user_api_key_auth(request_route="/v1/chat/completions"),
response=_async_chunk_iter(chunks),
request_data=data,
)
]
assert [chunk.model_dump() for chunk in delivered] == [chunk.model_dump() for chunk in chunks]
assert any(
"RaisingAcloseCallback" in message and "RuntimeError" in message and "cleanup failed" not in message
for message in _warnings(caplog)
)