diff --git a/litellm/llms/anthropic/chat/guardrail_translation/handler.py b/litellm/llms/anthropic/chat/guardrail_translation/handler.py index 806240c9749..578a4e47dfd 100644 --- a/litellm/llms/anthropic/chat/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/chat/guardrail_translation/handler.py @@ -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( diff --git a/litellm/llms/base_llm/guardrail_translation/base_translation.py b/litellm/llms/base_llm/guardrail_translation/base_translation.py index 89ad67f0485..70b8291d32c 100644 --- a/litellm/llms/base_llm/guardrail_translation/base_translation.py +++ b/litellm/llms/base_llm/guardrail_translation/base_translation.py @@ -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", diff --git a/litellm/llms/openai/chat/guardrail_translation/handler.py b/litellm/llms/openai/chat/guardrail_translation/handler.py index aa175733582..aee439862db 100644 --- a/litellm/llms/openai/chat/guardrail_translation/handler.py +++ b/litellm/llms/openai/chat/guardrail_translation/handler.py @@ -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]: diff --git a/litellm/llms/openai/responses/guardrail_translation/handler.py b/litellm/llms/openai/responses/guardrail_translation/handler.py index 90cdef87ec7..ad68924ce20 100644 --- a/litellm/llms/openai/responses/guardrail_translation/handler.py +++ b/litellm/llms/openai/responses/guardrail_translation/handler.py @@ -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( diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index c85169f0ba5..687df9b0348 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -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, diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 620b24df95d..f900d14bdc1 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index 37c1829def4..0c4ea6b29b5 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -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 diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index af487cb17a3..ed2a5b89a32 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 1c738aefb97..29f2f46f001 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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") diff --git a/tests/integration/observability/test_guardrail_effects.py b/tests/integration/observability/test_guardrail_effects.py index d377afb206c..1df4427bb4e 100644 --- a/tests/integration/observability/test_guardrail_effects.py +++ b/tests/integration/observability/test_guardrail_effects.py @@ -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 diff --git a/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py b/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py index 09921a9a85c..85e4c055f07 100644 --- a/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py +++ b/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py @@ -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 diff --git a/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py b/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py index 5c85faa5e13..5d71cbe9afc 100644 --- a/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py +++ b/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py @@ -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)) diff --git a/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py b/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py index 87980a47f87..8bd242308f7 100644 --- a/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py +++ b/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py @@ -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 diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py b/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py index 45ad336368c..0cae9344160 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py @@ -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( diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py b/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py index e547575ef9c..e797304a1c6 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py @@ -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", diff --git a/tests/unit/proxy/hooks/test_async_post_call_streaming_iterator_hook.py b/tests/unit/proxy/hooks/test_async_post_call_streaming_iterator_hook.py index 9c785d59830..31f1c054c59 100644 --- a/tests/unit/proxy/hooks/test_async_post_call_streaming_iterator_hook.py +++ b/tests/unit/proxy/hooks/test_async_post_call_streaming_iterator_hook.py @@ -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() + ] diff --git a/tests/unit/proxy/test_common_request_processing.py b/tests/unit/proxy/test_common_request_processing.py index 8485c286a30..8dc6c0477b0 100644 --- a/tests/unit/proxy/test_common_request_processing.py +++ b/tests/unit/proxy/test_common_request_processing.py @@ -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, ): diff --git a/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py b/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py index 2397c6480d5..e2fbd5dcd98 100644 --- a/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py +++ b/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py @@ -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(): """ diff --git a/tests/unit/proxy/utils/proxy_logging/test_guardrail_pipeline.py b/tests/unit/proxy/utils/proxy_logging/test_guardrail_pipeline.py index 8bd9dc0df8a..dbb354526d8 100644 --- a/tests/unit/proxy/utils/proxy_logging/test_guardrail_pipeline.py +++ b/tests/unit/proxy/utils/proxy_logging/test_guardrail_pipeline.py @@ -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) + )