mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(guardrails): run the end-of-stream post_call scan when the client disconnects mid-stream (#43839)
* fix(guardrails): run end-of-stream post_call scan when the client disconnects mid-stream Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): close the guardrail stream chain in async_data_generator on client disconnect Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): leave the raw upstream response to the shielded finalizer on client disconnect Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): keep disconnect cleanup going when a streaming callback cleanup raises Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): assert the refund through a recorder instead of the mock Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(guardrails): inspect tool calls released before a disconnect under incremental_diff and record a failed scan marker The incremental_diff transform stream now scans tool calls it already released when the client disconnects, and a disconnect scan whose translation raises after the guardrail recorded success also records guardrail_failed_to_respond Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(guardrails): pin that a text-only disconnect scan is not handed a tool_calls finish Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(guardrails): cover disconnect scans on every streaming endpoint and client, plus outage, worker-kill and cache-hit cells Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(guardrails): prove the cache-hit twin is served from the cache Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(guardrails): type the disconnect-close streams so basedpyright stops reporting unknown arguments * fix(guardrails): give the guardrail metadata cast a reason so the type discipline gate accepts it Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(guardrails): scan released Messages and Responses tool calls on disconnect Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * style(guardrails): format the disconnect scan unit tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore(guardrails): drop mutable-ok markers that no longer suppress a rule Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(guardrails): scan released Responses output after a finished item and end only in-flight Chat choices Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(guardrails): type the disconnect scan test helpers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(guardrails): type request_data in the disconnect scan helpers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(guardrails): type request_data in the disconnect scan test doubles Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(guardrails): pin that chat streams with no tool call in flight end as released Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(guardrails): scan only released chunks on disconnect and skip it once a block owns the verdict The disconnect scan now uses the chunks actually yielded to the client, copies them before scanning, skips when a mid-stream block or HTTP error already settled the verdict, and the iterator wrapper only closes hooks that are async generators so plain async iterator hooks keep working Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(guardrails): pin that a delivered guardrail error or final chunk settles the disconnect verdict Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(guardrails): close any hook iterator that exposes aclose when the stream ends early Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): accept a synchronous aclose on custom streaming hook iterators Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): swallow callback aclose errors at end of stream A custom callback whose async_post_call_streaming_iterator_hook returns a non-generator async iterator with a raising aclose() failed the finished stream: content plus usage reached the client and then the stream surfaced an error SSE with no [DONE], or aborted a post_call pipeline's buffering loop into a 500 with an empty body. Wrap the aclose invocation in _wrap_streaming_iterator_with_enrichment in try/except and log a warning naming the callback and the cleanup error, matching close_guarded_stream and _close_guarded_layers. Iteration-time hook exceptions still propagate. * fix(proxy): log only the error type when a callback aclose raises Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yucheng <yucheng@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
1238cfe90f
commit
7a7d27c550
19 changed files with 2527 additions and 229 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue