diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py index de389d8a945..320d687b3aa 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py @@ -39,6 +39,9 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" streaming_end_of_stream_only=_get_config_value(litellm_params, optional_params, "streaming_end_of_stream_only"), streaming_sampling_rate=_get_config_value(litellm_params, optional_params, "streaming_sampling_rate"), streaming_transform_mode=_get_config_value(litellm_params, optional_params, "streaming_transform_mode"), + streaming_buffer_until_moderated=_get_config_value( + litellm_params, optional_params, "streaming_buffer_until_moderated" + ), timeout=litellm_params.timeout, ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py index 7bb41b7586b..c745a6c91b3 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py @@ -204,6 +204,7 @@ class GenericGuardrailAPI(CustomGuardrail): streaming_end_of_stream_only: bool | None = None, streaming_sampling_rate: int | None = None, streaming_transform_mode: Literal["block_only", "incremental_diff"] | None = None, + streaming_buffer_until_moderated: bool | None = None, **kwargs, ): self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback) @@ -250,6 +251,8 @@ class GenericGuardrailAPI(CustomGuardrail): self.streaming_transform_mode: Literal["block_only", "incremental_diff"] = ( "block_only" if streaming_transform_mode is None else streaming_transform_mode ) + if streaming_buffer_until_moderated is not None: + self.streaming_buffer_until_moderated: bool = streaming_buffer_until_moderated # Set supported event hooks kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks())) diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 196e0f20ad3..0938d3f35c3 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -100,6 +100,18 @@ def _proxy_stamped_used_client_oauth_token( return stamped if isinstance(stamped, bool) else None +def _failure_snapshot_with_guardrails_recorded_after_it( + failure_snapshot: object, recorded_guardrail_information: object +) -> Mapping[str, object] | None: + if ( + not recorded_guardrail_information + or not isinstance(failure_snapshot, dict) + or failure_snapshot.get("guardrail_information") + ): + return None + return {**failure_snapshot, "guardrail_information": recorded_guardrail_information} + + def _proxy_spend_writer() -> DBSpendUpdateWriter: from litellm.proxy.proxy_server import proxy_logging_obj @@ -275,13 +287,22 @@ class _ProxyDBLogger(CustomLogger): existing_metadata.get("standard_logging_guardrail_information") ) + failure_snapshot_with_late_guardrails: Final = _failure_snapshot_with_guardrails_recorded_after_it( + request_data.get("standard_logging_object"), + existing_metadata.get("standard_logging_guardrail_information"), + ) + await self._spend_writer().update_database( token=LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict), response_cost=recovered_response_cost, user_id=user_api_key_dict.user_id, end_user_id=user_api_key_dict.end_user_id, team_id=user_api_key_dict.team_id, - kwargs=request_data, + kwargs=( + request_data + if failure_snapshot_with_late_guardrails is None + else {**request_data, "standard_logging_object": failure_snapshot_with_late_guardrails} + ), completion_response=original_exception, start_time=actual_start_time, end_time=datetime.now(), diff --git a/litellm/proxy/policy_engine/pipeline_executor.py b/litellm/proxy/policy_engine/pipeline_executor.py index 5ccf2920030..f47b41592c2 100644 --- a/litellm/proxy/policy_engine/pipeline_executor.py +++ b/litellm/proxy/policy_engine/pipeline_executor.py @@ -32,6 +32,7 @@ from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrai ) from litellm.types.guardrails import GuardrailEventHooks from litellm.types.proxy.policy_engine.pipeline_types import ( + DETECT_ONLY_PIPELINE_ACTIONS, PipelineExecutionResult, PipelineStep, PipelineStepResult, @@ -119,9 +120,9 @@ class _StreamRewriteObserver(CustomGuardrail): guardrail. It records whether the guardrail returned different output than it was given, which for guardrails like Bedrock's ANONYMIZED action is only known at runtime. Text and tool-call rewrites are deliverable on translations that write them back across the - buffered chunks (``delivers_ended_stream_rewrites``); rewrites on any other translation, - and a rewrite that drops or adds a tool call on any translation, are discarded by the - executor, which releases the original chunks. + buffered chunks (``delivers_ended_stream_rewrites``); rewrites on any other translation or + on a stream the client already received live, and a rewrite that drops or adds a tool + call on any translation, are discarded by the executor, which releases the original chunks. The inner guardrail's ``apply_guardrail`` already records the guardrail information and span, so the observer's stays out of ``log_guardrail_information``.""" @@ -156,18 +157,26 @@ class _StreamRewriteObserver(CustomGuardrail): ) return outputs - def discard_reason(self, deliver_rewrites: bool) -> str | None: + def discard_reason(self, rewrite_undeliverable_reason: str | None) -> str | None: if self.tool_call_count_change is not None: sent, returned = self.tool_call_count_change return ( f"the guardrail returned {returned} tool calls for a stream that carried {sent}, and a rewrite " "that drops or adds a tool call cannot be written back" ) - if not deliver_rewrites and (self.rewrote_texts or self.rewrote_tool_calls): - return "this endpoint's streaming pipeline does not write ended-stream rewrites back yet" + if rewrite_undeliverable_reason is not None and (self.rewrote_texts or self.rewrote_tool_calls): + return rewrite_undeliverable_reason return None +def _rewrite_undeliverable_reason(translation_delivers_rewrites: bool, stream_already_sent: bool) -> str | None: + if stream_already_sent: + return "the client already received the stream live" + if not translation_delivers_rewrites: + return "this endpoint's streaming pipeline does not write ended-stream rewrites back yet" + return None + + class _ScannedTextRecorder(CustomGuardrail): def __init__(self, guardrail_name: str) -> None: super().__init__(guardrail_name=guardrail_name) @@ -296,7 +305,7 @@ def _prepare_hook_input( callback: CustomGuardrail, data: dict, # mutable-ok: same request-payload shape the hooks mutate raw_request_snapshot: dict | None, # mutable-ok: same request-payload shape as data -) -> tuple[dict, bool]: # mutable-ok: returns that same request-payload dict +) -> tuple[dict[str, object], bool]: # mutable-ok: returns that same request-payload dict """Inject the step's guardrail name into metadata so should_run_guardrail() allows it, and pick the payload the step scans: a scan_raw_request step evaluates the pristine pre-pipeline snapshot instead of `data` (which earlier pass_data steps in this same @@ -344,6 +353,7 @@ class PipelineExecutor: raw_request_snapshot: dict | None = None, # mutable-ok: same request-payload shape as data streaming_chunks: list[object] | None = None, # mutable-ok: shared buffered-stream chunks, read per step endpoint_translation: "BaseTranslation | None" = None, + stream_already_sent: bool = False, ) -> PipelineExecutionResult: """ Execute pipeline steps sequentially with conditional actions. @@ -366,6 +376,9 @@ class PipelineExecutor: instead of calling ``async_post_call_success_hook``. endpoint_translation: the guardrail translation for the streamed endpoint, resolved by the caller. + stream_already_sent: the client already received ``streaming_chunks`` + live, so each step's rewrite is discarded right after the step and + every later step scans what the client received. Returns: PipelineExecutionResult with terminal action and step results @@ -392,6 +405,7 @@ class PipelineExecutor: raw_request_snapshot=raw_request_snapshot, streaming_chunks=streaming_chunks, endpoint_translation=endpoint_translation, + stream_already_sent=stream_already_sent, ) duration = time.perf_counter() - start_time @@ -460,13 +474,15 @@ class PipelineExecutor: hook_input: dict[str, object], # mutable-ok: same request-payload shape as data user_api_key_dict: "UserAPIKeyAuth", litellm_logging_obj: "LiteLLMLoggingObj | None", + stream_already_sent: bool, ) -> None: """Run one streaming post_call step through the endpoint translation, delivering text and tool-call rewrites on translations that support ended-stream write-back. A guardrail without the unified interface runs its legacy post-call hook against the assembled response through ``_LegacyHookStreamAdapter``. A rewrite that cannot reach the - client yet (one on a translation without write-back, one that drops or adds a tool call, - or one the translation or adapter refused with ``UndeliverableStreamRewrite``) is + client (one on a translation without write-back, one on a stream the client already received + live, one that drops or adds a tool call, or one the translation or adapter refused with + ``UndeliverableStreamRewrite``) is discarded: the buffered chunks go back to the originals and the step passes, so the client gets the stream the merge base sent, and the guardrail stays out of the applied-guardrails header since its output never reached the client. The response an @@ -478,11 +494,13 @@ class PipelineExecutor: else _LegacyHookStreamAdapter(callback, endpoint_translation, user_api_key_dict) ) observer: Final = _StreamRewriteObserver(scanner) - deliver_rewrites: Final = type(endpoint_translation).delivers_ended_stream_rewrites + rewrite_undeliverable_reason: Final = _rewrite_undeliverable_reason( + type(endpoint_translation).delivers_ended_stream_rewrites, stream_already_sent + ) originals: Final = copy.deepcopy(streaming_chunks) hook_input.pop("response", None) try: - if deliver_rewrites: + if rewrite_undeliverable_reason is None: await endpoint_translation.process_output_streaming_response( responses_so_far=streaming_chunks, guardrail_to_apply=observer, @@ -502,7 +520,7 @@ class PipelineExecutor: except UndeliverableStreamRewrite as undeliverable: _release_original_chunks(step.guardrail, undeliverable.reason, streaming_chunks, originals) return - discard_reason: Final = observer.discard_reason(deliver_rewrites) + discard_reason: Final = observer.discard_reason(rewrite_undeliverable_reason) if discard_reason is not None: _release_original_chunks(step.guardrail, discard_reason, streaming_chunks, originals) return @@ -519,6 +537,7 @@ class PipelineExecutor: raw_request_snapshot: dict | None = None, # mutable-ok: same request-payload shape as data streaming_chunks: list[object] | None = None, # mutable-ok: shared buffered-stream chunks, read per step endpoint_translation: "BaseTranslation | None" = None, + stream_already_sent: bool = False, ) -> tuple[ Literal["pass", "fail", "error", "skip"], dict | None, @@ -547,7 +566,7 @@ class PipelineExecutor: return ("skip", None, None, None) hook_input, scans_raw_request = _prepare_hook_input(step, callback, data, raw_request_snapshot) - snapshot_entries_before: Final = len(_recorded_guardrail_information(hook_input)) + snapshot_entries_before: Final = len(recorded_guardrail_information(hook_input)) # Use unified_guardrail path if callback implements apply_guardrail target: CustomLogger = callback @@ -584,6 +603,7 @@ class PipelineExecutor: hook_input=hook_input, user_api_key_dict=user_api_key_dict, litellm_logging_obj=data.get("litellm_logging_obj"), + stream_already_sent=stream_already_sent, ) response = None elif mode == "post_call": @@ -624,7 +644,7 @@ class PipelineExecutor: if hook_input is not data: _append_guardrail_information( request_data=data, - entries=_recorded_guardrail_information(hook_input)[snapshot_entries_before:], + entries=recorded_guardrail_information(hook_input)[snapshot_entries_before:], ) @staticmethod @@ -697,7 +717,7 @@ def _restore_request_guardrails( _GUARDRAIL_INFORMATION_KEY: Final = "standard_logging_guardrail_information" -def _recorded_guardrail_information(source: Mapping[str, object]) -> list[StandardLoggingGuardrailInformation]: +def recorded_guardrail_information(source: Mapping[str, object]) -> list[StandardLoggingGuardrailInformation]: bucket: Final = source.get(get_metadata_variable_name_from_kwargs(source)) recorded: Final = bucket.get(_GUARDRAIL_INFORMATION_KEY) if isinstance(bucket, dict) else None return recorded if isinstance(recorded, list) else [] @@ -721,8 +741,8 @@ def _carry_working_guardrail_information( working_data: Mapping[str, object], request_data: dict[str, object], # mutable-ok: same request-payload shape as execute_steps' data ) -> None: - recorded: Final = _recorded_guardrail_information(working_data) - existing: Final = _recorded_guardrail_information(request_data) + recorded: Final = recorded_guardrail_information(working_data) + existing: Final = recorded_guardrail_information(request_data) if recorded is existing: return _append_guardrail_information(request_data=request_data, entries=[e for e in recorded if e not in existing]) @@ -748,6 +768,16 @@ def _pipeline_action_for_outcome(step: PipelineStep, outcome: str) -> str: return step.on_fail +def pipeline_step_is_detect_only(step: PipelineStep) -> bool: + """Whether every action the step can take leaves the response as the provider sent it. A block or a + modify_response acts on the stream, so a step that can reach either must see its verdict before the + client sees the chunks; a step that can only allow or pass to the next step takes its verdict after""" + reachable_actions: Final = frozenset( + _pipeline_action_for_outcome(step, outcome) for outcome in ("pass", "fail", "error") + ) + return reachable_actions <= DETECT_ONLY_PIPELINE_ACTIONS + + def _extract_error_message(e: Exception) -> str: """Extract a human-readable error message from a guardrail exception.""" if isinstance(e, ModifyResponseException): diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 8f4bd9404de..302f55eadf4 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -19,6 +19,7 @@ from collections.abc import ( Awaitable, Callable, Coroutine, + Iterator, Mapping, Sequence, ) @@ -47,6 +48,7 @@ from typing import ( runtime_checkable, ) +import anyio from typing_extensions import Never, ReadOnly, TypedDict from litellm import _custom_logger_compatible_callbacks_literal @@ -236,7 +238,11 @@ from litellm.proxy.hooks.sensitive_data_routing import ( # noqa: F401, RUF100 ) from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup, add_guardrails_from_auth_metadata from litellm.proxy.management_helpers.key_settings_audit import with_settings_updated_at -from litellm.proxy.policy_engine.pipeline_executor import PipelineExecutor +from litellm.proxy.policy_engine.pipeline_executor import ( + PipelineExecutor, + pipeline_step_is_detect_only, + recorded_guardrail_information, +) from litellm.proxy.policy_engine.policy_registry import get_policy_registry from litellm.proxy.policy_engine.policy_resolver import PolicyResolver from litellm.repositories.budget_repository import BudgetRepository @@ -286,7 +292,7 @@ if TYPE_CHECKING: from litellm.proxy.db.spend_log_tool_index import ToolUsageTransaction from litellm.repositories.prisma_protocols import TableActions from litellm.types.mcp_server.mcp_server_manager import MCPServer - from litellm.types.proxy.policy_engine.pipeline_types import GuardrailPipeline + from litellm.types.proxy.policy_engine.pipeline_types import GuardrailPipeline, PipelineStep Span = _Span | object else: @@ -597,8 +603,30 @@ def _policy_pipelines(data: Mapping[str, object]) -> tuple[tuple[str, "Guardrail ) +def _pipeline_steps(pipelines: Sequence[tuple[str, "GuardrailPipeline"]]) -> Iterator["PipelineStep"]: + for _policy_name, pipeline in pipelines: + yield from pipeline.steps + + def _pipeline_step_guardrail_names(pipelines: Sequence[tuple[str, "GuardrailPipeline"]]) -> frozenset[str]: - return frozenset(step.guardrail for _policy_name, pipeline in pipelines for step in pipeline.steps) + return frozenset(step.guardrail for step in _pipeline_steps(pipelines)) + + +def _pipeline_step_streams_live(step: "PipelineStep") -> bool: + """A detect-only step whose guardrail set ``streaming_buffer_until_moderated: false`` takes its verdict after + the chunks reach the client. A guardrail that masks response content keeps buffering, since the client would + already hold the unredacted stream, and a guardrail that is not loaded keeps buffering so the executor's + not-found outcome settles the stream before release""" + if not pipeline_step_is_detect_only(step): + return False + callback: Final = PipelineExecutor.find_guardrail_callback(step.guardrail) + if callback is None or callback.mask_response_content: + return False + return not unified_guardrail.resolve_streaming_flag(callback, "streaming_buffer_until_moderated", True) + + +def _pipelines_stream_live(pipelines: Sequence[tuple[str, "GuardrailPipeline"]]) -> bool: + return all(_pipeline_step_streams_live(step) for step in _pipeline_steps(pipelines)) def pipeline_managed_guardrail_names( @@ -2287,10 +2315,10 @@ class ProxyLogging: @staticmethod def _handle_pipeline_result( result: PipelineExecutionResult, - data: dict, + data: dict[str, object], policy_name: str, original_response: "LLMResponseTypes | Sequence[object] | None" = None, - ) -> dict: + ) -> dict[str, object]: """ Handle a PipelineExecutionResult — allow, block, or modify_response. @@ -2352,9 +2380,10 @@ class ProxyLogging: raise HTTPException(status_code=400, detail=error_detail) if result.terminal_action == "modify_response": + model: Final = data.get("model") raise ModifyResponseException( message=result.modify_response_message or "Response modified by pipeline", - model=data.get("model", "unknown"), + model=model if isinstance(model, str) else "unknown", request_data=data, guardrail_name=f"pipeline:{policy_name}", detection_info=None, @@ -4059,7 +4088,12 @@ class ProxyLogging: resolve_endpoint_translation(user_api_key_dict, None) if post_call_pipelines else None ) if pipeline_translation is not None: - current_response = self._pipeline_gated_stream( + pipeline_stream: Final = ( + self._pipeline_scanned_live_stream + if _pipelines_stream_live(post_call_pipelines) + else self._pipeline_gated_stream + ) + current_response = pipeline_stream( response=current_response, user_api_key_dict=user_api_key_dict, request_data=request_data, @@ -4094,7 +4128,7 @@ class ProxyLogging: self, response: "AsyncGenerator[object, None]", user_api_key_dict: UserAPIKeyAuth, - request_data: dict, # mutable-ok: same request-payload shape the hooks mutate + request_data: dict[str, object], # mutable-ok: same request-payload shape the hooks mutate pipelines: "tuple[tuple[str, GuardrailPipeline], ...]", translation: "tuple[str, BaseTranslation]", ) -> "AsyncGenerator[object, None]": @@ -4155,6 +4189,125 @@ class ProxyLogging: for buffered_item in buffered: yield buffered_item + async def _pipeline_scanned_live_stream( + self, + response: "AsyncGenerator[object, None]", + user_api_key_dict: UserAPIKeyAuth, + request_data: dict[str, object], # mutable-ok: same request-payload shape the hooks mutate + pipelines: "tuple[tuple[str, GuardrailPipeline], ...]", + translation: "tuple[str, BaseTranslation]", + ) -> "AsyncGenerator[object, None]": + """ + Execute detect-only post_call policy pipelines against a streamed response the client receives live. + + Every chunk reaches the client as the provider sends it. Once the stream ends, each pipeline's steps run + against a copy of the assembled output, the same end-of-stream scan ``_pipeline_gated_stream`` runs + before release. The steps can only allow or pass to the next one, so the scan records each guardrail's + verdict in guardrail_information without changing what was sent; the executor discards a rewrite right + after the step that returned it, with a warning, so every step scans the text the client received. A + stream cut short, by a client disconnect or by a provider error after some chunks went out, still gets + the scan over what the client received, shielded from cancellation, and a scan that raises then records + guardrail_failed_to_respond for every step guardrail without a verdict, so the spend log never shows a + silent skip. + """ + released: Final[list[object]] = [] # mutable-ok: accumulates the chunks the client already received + try: + async for item in response: + released.append(item) + yield item + except (GeneratorExit, asyncio.CancelledError, Exception): + if released: + await self._scan_pipeline_stream_cut_short( + released=released, + user_api_key_dict=user_api_key_dict, + request_data=request_data, + pipelines=pipelines, + translation=translation, + ) + raise + if not released: + return + with anyio.CancelScope(shield=True): + await self._scan_released_pipeline_stream( + originals=released, + user_api_key_dict=user_api_key_dict, + request_data=request_data, + pipelines=pipelines, + translation=translation, + ) + + async def _scan_released_pipeline_stream( + self, + *, + originals: Sequence[object], + user_api_key_dict: UserAPIKeyAuth, + request_data: dict[str, object], # mutable-ok: same request-payload shape the hooks mutate + pipelines: "tuple[tuple[str, GuardrailPipeline], ...]", + translation: "tuple[str, BaseTranslation]", + ) -> None: + call_type, endpoint_translation = translation + scanned: Final[list[object]] = list(copy.deepcopy(tuple(originals))) # mutable-ok: executor scans in place + for policy_name, pipeline in pipelines: + result: PipelineExecutionResult = await PipelineExecutor.execute_steps( + steps=pipeline.steps, + mode="post_call", + data=request_data, + user_api_key_dict=user_api_key_dict, + call_type=call_type, + policy_name=policy_name, + streaming_chunks=scanned, + endpoint_translation=endpoint_translation, + stream_already_sent=True, + ) + ProxyLogging._handle_pipeline_result( + result, data=request_data, policy_name=policy_name, original_response=originals + ) + + async def _scan_pipeline_stream_cut_short( + self, + *, + released: Sequence[object], + user_api_key_dict: UserAPIKeyAuth, + request_data: dict[str, object], # mutable-ok: same request-payload shape the hooks mutate + pipelines: "tuple[tuple[str, GuardrailPipeline], ...]", + translation: "tuple[str, BaseTranslation]", + ) -> None: + _call_type, endpoint_translation = translation + recorded_before: Final = len(recorded_guardrail_information(request_data)) + with anyio.CancelScope(shield=True): + try: + await self._scan_released_pipeline_stream( + originals=endpoint_translation.released_stream_as_ended(copy.deepcopy(tuple(released))), + user_api_key_dict=user_api_key_dict, + request_data=request_data, + pipelines=pipelines, + translation=translation, + ) + except Exception as e: # noqa: BLE001 # the stream already ended, so the verdict can only be recorded + verbose_proxy_logger.warning( + "Policy pipelines scanned a stream that was cut short and raised %s", + type(e).__name__, + ) + recorded_names: Final = frozenset( + entry["guardrail_name"] + for entry in recorded_guardrail_information(request_data)[recorded_before:] + if "guardrail_name" in entry + ) + unsettled: Final = tuple( + PipelineExecutor.find_guardrail_callback(step.guardrail) + for step in _pipeline_steps(pipelines) + if step.guardrail not in recorded_names + ) + for callback in unsettled: + if callback is None: + continue + callback.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, + ) + @staticmethod async def _close_guarded_layers(layers: Sequence[AsyncGenerator[object, None]]) -> None: for layer in reversed(layers): diff --git a/litellm/types/proxy/policy_engine/pipeline_types.py b/litellm/types/proxy/policy_engine/pipeline_types.py index 0a089d2fe9b..98b03b0dd79 100644 --- a/litellm/types/proxy/policy_engine/pipeline_types.py +++ b/litellm/types/proxy/policy_engine/pipeline_types.py @@ -13,6 +13,7 @@ from pydantic import ConfigDict, Field, field_validator from litellm.types.llms.base import LiteLLMBaseModel VALID_PIPELINE_ACTIONS: Final = {"allow", "block", "next", "modify_response"} +DETECT_ONLY_PIPELINE_ACTIONS: Final = frozenset({"allow", "next"}) VALID_PIPELINE_MODES: Final = {"pre_call", "post_call"} diff --git a/tests/integration/observability/test_guardrail_effects.py b/tests/integration/observability/test_guardrail_effects.py index 79feb4f0d92..5f1219e18cb 100644 --- a/tests/integration/observability/test_guardrail_effects.py +++ b/tests/integration/observability/test_guardrail_effects.py @@ -2159,8 +2159,21 @@ def _failing_response_scans(reply: Reply) -> Callable[[Request], Reply]: return guardrail +def _detect_only_pipeline_policy(identity: str) -> dict[str, JsonValue]: + return { + "guardrails": {"add": [identity]}, + "pipeline": {"mode": "post_call", "steps": [{"guardrail": identity, "on_pass": "allow", "on_fail": "next"}]}, + } + + def _post_call_config( - tmp_path: Path, identity: str, policy_url: str, params: Mapping[str, JsonValue], default_on: bool + tmp_path: Path, + identity: str, + policy_url: str, + params: Mapping[str, JsonValue], + default_on: bool, + *, + pipeline: bool = False, ) -> Path: config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) config["guardrails"] = [ @@ -2176,6 +2189,9 @@ def _post_call_config( }, } ] + if pipeline: + config["policies"] = {f"{identity}-pipeline": _detect_only_pipeline_policy(identity)} + config["policy_attachments"] = [{"policy": f"{identity}-pipeline", "scope": "*"}] path: Final = tmp_path / f"{identity}.yaml" path.write_text(yaml.safe_dump(config)) return path @@ -2227,6 +2243,7 @@ def _disconnect_rig( shape: str = "text", default_on: bool = True, workers: int = 1, + pipeline: bool = False, ) -> Iterator[_DisconnectRig]: identity: Final = "guardrail" + uuid.uuid4().hex gate: Final = threading.Event() @@ -2234,7 +2251,7 @@ def _disconnect_rig( 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) + config: Final = _post_call_config(tmp_path, identity, policy.url, params, default_on, pipeline=pipeline) try: with ( owned_proxy_process(gateway, tmp_path, {}, config=config, workers=workers) as owned, @@ -2602,6 +2619,37 @@ def test_client_disconnect_mid_stream_scans_for_a_guardrail_the_request_opted_in assert _post_call_statuses(rig, rig.model) == (("success",),) +_LIVE_DETECT_ONLY_PIPELINE: Final = MappingProxyType({"streaming_buffer_until_moderated": False}) + + +@pytest.mark.parametrize( + "disconnect", + ( + pytest.param(_chat_httpx, id="chat-httpx"), + pytest.param(_responses_httpx, id="responses-httpx"), + pytest.param(_messages_httpx, id="messages-httpx"), + ), +) +def test_live_detect_only_pipeline_releases_chunks_before_the_scan_and_scans_what_a_disconnected_client_received( + gateway: Gateway, tmp_path: Path, disconnect: Callable[[_DisconnectRig, int], str] +) -> None: + with _disconnect_rig( + gateway, tmp_path, params=_LIVE_DETECT_ONLY_PIPELINE, default_on=False, pipeline=True + ) as rig: + scans: Final = _scanned_while_upstream_is_held(rig, disconnect) + assert len(scans) == 1, scans + + +def test_live_detect_only_pipeline_records_the_verdict_of_a_disconnected_stream_on_the_spend_row( + gateway: Gateway, tmp_path: Path +) -> None: + with _disconnect_rig( + gateway, tmp_path, params=_LIVE_DETECT_ONLY_PIPELINE, default_on=False, pipeline=True + ) as rig: + _scanned_while_upstream_is_held(rig, _chat_httpx) + 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: @@ -2614,16 +2662,17 @@ def test_client_disconnect_before_any_content_sends_no_response_scan(gateway: Ga @pytest.mark.parametrize( - "params", + ("params", "pipeline"), ( - pytest.param(dict(_END_OF_STREAM_ONLY), id="end-of-stream-only"), - pytest.param({"streaming_buffer_until_moderated": True}, id="buffered"), + pytest.param(dict(_END_OF_STREAM_ONLY), False, id="end-of-stream-only"), + pytest.param({"streaming_buffer_until_moderated": True}, False, id="buffered"), + pytest.param(dict(_LIVE_DETECT_ONLY_PIPELINE), True, id="live-detect-only-pipeline"), ), ) def test_a_fully_read_stream_is_scanned_exactly_once( - gateway: Gateway, tmp_path: Path, params: dict[str, JsonValue] + gateway: Gateway, tmp_path: Path, params: dict[str, JsonValue], pipeline: bool ) -> None: - with _disconnect_rig(gateway, tmp_path, params=params, gated=False) as rig: + with _disconnect_rig(gateway, tmp_path, params=params, gated=False, default_on=not pipeline, pipeline=pipeline) 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 diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py b/tests/unit/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py index e97de4686bf..7d70b96eeff 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py @@ -6,6 +6,7 @@ specifically focusing on metadata extraction and passing. """ import os +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import httpx @@ -1332,6 +1333,44 @@ class TestGenericGuardrailAPIStreamingConfig: assert guardrail.streaming_end_of_stream_only is False assert guardrail.streaming_sampling_rate == 3 + @pytest.mark.parametrize( + ("configured", "streams_live"), + [ + pytest.param(None, False, id="unset-keeps-the-pipeline-buffered"), + pytest.param(True, False, id="true-keeps-the-pipeline-buffered"), + pytest.param(False, True, id="false-streams-the-pipeline-live"), + ], + ) + def test_initialize_guardrail_streaming_buffer_until_moderated_reaches_the_pipeline_live_check( + self, monkeypatch: pytest.MonkeyPatch, configured: bool | None, streams_live: bool + ) -> None: + from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import ( + initialize_guardrail, + ) + from litellm.proxy.utils import _pipelines_stream_live + from litellm.types.guardrails import LitellmParams + from litellm.types.proxy.policy_engine.pipeline_types import GuardrailPipeline, PipelineStep + + litellm_params: Final = LitellmParams.model_validate( + { + "guardrail": "generic_guardrail_api", + "mode": "post_call", + "api_base": "https://api.test.guardrail.com", + "default_on": False, + **({} if configured is None else {"streaming_buffer_until_moderated": configured}), + } + ) + + with patch("litellm.logging_callback_manager.add_litellm_callback"): + guardrail: Final = initialize_guardrail(litellm_params, {"guardrail_name": "pipeline-scanner"}) + monkeypatch.setattr(litellm, "callbacks", [guardrail]) + pipeline: Final = GuardrailPipeline( + mode="post_call", + steps=[PipelineStep(guardrail="pipeline-scanner", on_pass="allow", on_fail="next")], + ) + + assert _pipelines_stream_live((("detect-only", pipeline),)) is streams_live + def test_initialize_guardrail_optional_params_defaults_do_not_shadow_top_level( self, ): diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrails/test_streaming_buffer_until_moderated.py b/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrails/test_streaming_buffer_until_moderated.py index db937f18e96..2436c85fedd 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrails/test_streaming_buffer_until_moderated.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrails/test_streaming_buffer_until_moderated.py @@ -15,6 +15,7 @@ from typing import Any, AsyncGenerator, List, Literal, Optional import pytest +import litellm from litellm.integrations.custom_guardrail import ( CustomGuardrail, ModifyResponseException, @@ -25,6 +26,8 @@ from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrai UnifiedLLMGuardrails, _is_redundant_scan, ) +from litellm.proxy.utils import _pipelines_stream_live +from litellm.types.proxy.policy_engine.pipeline_types import GuardrailPipeline, PipelineStep from litellm.types.utils import ( ChatCompletionDeltaToolCall, Delta, @@ -571,3 +574,34 @@ async def test_buffered_mode_disabled_for_content_rewriting_guardrail(): assert guardrail.streaming_buffer_until_moderated is True # request asked for buffering assert ORIGINAL_MARKER in raw assert BLOCK_MESSAGE not in raw + + +def _detect_only_pipeline(guardrail: CustomGuardrail) -> tuple[tuple[str, GuardrailPipeline], ...]: + step = PipelineStep(guardrail=guardrail.guardrail_name, on_pass="allow", on_fail="next") + return (("live-policy", GuardrailPipeline(mode="post_call", steps=[step])),) + + +@pytest.mark.parametrize( + ("attribute", "guardrail_config", "streams_live"), + ( + pytest.param(False, None, True, id="attribute-false-streams-live"), + pytest.param(None, None, False, id="unset-keeps-buffering"), + pytest.param(True, None, False, id="attribute-true-keeps-buffering"), + pytest.param(False, {"streaming_buffer_until_moderated": True}, False, id="config-true-wins"), + pytest.param(True, {"streaming_buffer_until_moderated": False}, True, id="config-false-wins"), + ), +) +def test_pipeline_live_streaming_resolves_the_flag_like_the_flat_streaming_path( + monkeypatch: pytest.MonkeyPatch, + attribute: bool | None, + guardrail_config: dict | None, + streams_live: bool, +) -> None: + guardrail = _PassingGuardrail(guardrail_name="pipeline-scanner", event_hook="post_call") + if attribute is not None: + guardrail.streaming_buffer_until_moderated = attribute + if guardrail_config is not None: + guardrail.guardrail_config = guardrail_config + monkeypatch.setattr(litellm, "callbacks", [guardrail]) + + assert _pipelines_stream_live(_detect_only_pipeline(guardrail)) is streams_live diff --git a/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py b/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py index 811176c4f4e..100e79e371d 100644 --- a/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py @@ -2,6 +2,7 @@ import asyncio import json import logging from datetime import datetime, timedelta, timezone +from types import SimpleNamespace from typing import Final from unittest.mock import AsyncMock, MagicMock, patch @@ -356,6 +357,61 @@ async def test_async_post_call_failure_hook_adds_guardrail_cost_to_recovered_str assert mock_update_database.call_args[1]["response_cost"] == pytest.approx(0.0013) +@pytest.mark.asyncio +@pytest.mark.parametrize( + "snapshot_guardrail_information, expected_guardrail_names", + [ + (None, ["detect-only-scanner"]), + ([], ["detect-only-scanner"]), + ([{"guardrail_name": "logged-before-the-failure", "guardrail_status": "success"}], ["logged-before-the-failure"]), + ], +) +async def test_async_post_call_failure_hook_keeps_a_guardrail_verdict_recorded_after_the_failure_was_logged( + snapshot_guardrail_information: list[dict[str, str]] | None, expected_guardrail_names: list[str] +): + """ + A provider that drops a live stream mid-flight gets the failure logged at once, and only then does the + post_call scan of the chunks the client already received record its verdict, so the failure row showed no + guardrail at all while billing the scan's cost + """ + writer: Final = MagicMock(spec=DBSpendUpdateWriter) + writer.update_database = AsyncMock() + logger: Final = ProxyDBLogger(spend_writer=lambda: writer) + failure_snapshot: Final = { + "metadata": {}, + "hidden_params": {}, + "model_map_information": {}, + "guardrail_information": snapshot_guardrail_information, + } + + await logger.async_post_call_failure_hook( + request_data={ + "model": "gpt-5.6-sol", + "messages": [{"role": "user", "content": "Hello"}], + "metadata": { + "standard_logging_guardrail_information": [ + {"guardrail_name": "detect-only-scanner", "guardrail_status": "success"} + ] + }, + "litellm_logging_obj": SimpleNamespace( + model_call_details={"standard_logging_object": failure_snapshot}, litellm_trace_id="trace-1" + ), + }, + original_exception=Exception("provider dropped the stream"), + user_api_key_dict=UserAPIKeyAuth(api_key="test_api_key"), + ) + + row: Final = get_logging_payload( + kwargs=writer.update_database.call_args.kwargs["kwargs"], + response_obj={}, + start_time=datetime.now(), + end_time=datetime.now(), + ) + logged_guardrails: Final = json.loads(row["metadata"])["guardrail_information"] + assert [entry["guardrail_name"] for entry in logged_guardrails] == expected_guardrail_names + assert failure_snapshot["guardrail_information"] == snapshot_guardrail_information + + @pytest.mark.asyncio async def test_async_post_call_failure_hook_non_llm_route(): # Setup diff --git a/tests/unit/proxy/policy_engine/test_pipeline_executor.py b/tests/unit/proxy/policy_engine/test_pipeline_executor.py index d13abbd379b..8dbab66ff3d 100644 --- a/tests/unit/proxy/policy_engine/test_pipeline_executor.py +++ b/tests/unit/proxy/policy_engine/test_pipeline_executor.py @@ -19,7 +19,11 @@ from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.guardrails.guardrail_hooks.custom_code.custom_code_guardrail import ( CustomCodeGuardrail, ) -from litellm.proxy.policy_engine.pipeline_executor import PipelineExecutor, UndeliverableStreamRewrite +from litellm.proxy.policy_engine.pipeline_executor import ( + PipelineExecutor, + UndeliverableStreamRewrite, + pipeline_step_is_detect_only, +) from litellm.types.guardrails import GuardrailEventHooks from litellm.types.proxy.policy_engine.pipeline_types import ( GuardrailPipeline, @@ -1794,3 +1798,24 @@ def test_undeliverable_stream_rewrite_keeps_its_reason_through_a_copy(clone): assert copied.reason == "the translation refused it" assert str(copied) == str(original) assert str(copied).endswith("cannot be written back to the stream: the translation refused it") + + +@pytest.mark.parametrize( + ("on_pass", "on_fail", "on_error", "detect_only"), + ( + pytest.param("allow", "next", None, True, id="allow-next"), + pytest.param("next", "allow", "next", True, id="next-allow-next"), + pytest.param("allow", "next", "allow", True, id="allow-next-allow"), + pytest.param("allow", "block", None, False, id="on_fail-block"), + pytest.param("allow", "next", "block", False, id="on_error-block"), + pytest.param("allow", "modify_response", None, False, id="on_fail-modify_response"), + pytest.param("modify_response", "next", None, False, id="on_pass-modify_response"), + pytest.param("block", "next", None, False, id="on_pass-block"), + ), +) +def test_pipeline_step_is_detect_only_when_no_reachable_action_touches_the_response( + on_pass: str, on_fail: str, on_error: str | None, detect_only: bool +) -> None: + step = PipelineStep(guardrail="scanner", on_pass=on_pass, on_fail=on_fail, on_error=on_error) + + assert pipeline_step_is_detect_only(step) is detect_only diff --git a/tests/unit/proxy/utils/proxy_logging/test_guardrail_pipeline.py b/tests/unit/proxy/utils/proxy_logging/test_guardrail_pipeline.py index dbb354526d8..acc1463b072 100644 --- a/tests/unit/proxy/utils/proxy_logging/test_guardrail_pipeline.py +++ b/tests/unit/proxy/utils/proxy_logging/test_guardrail_pipeline.py @@ -14,7 +14,7 @@ import json from collections.abc import AsyncGenerator, Iterator from copy import deepcopy import logging -from typing import Any, Callable, Dict, List +from typing import Any, Callable, Dict, Final, List from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -31,6 +31,8 @@ from litellm.integrations.prometheus import PrometheusLogger from litellm.llms.base_llm.guardrail_translation.utils import stream_item_field from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.callback_utils import add_guardrail_to_applied_guardrails_header +from litellm.llms.openai.chat.guardrail_translation.handler import OpenAIChatCompletionsHandler +from litellm.proxy.policy_engine.pipeline_executor import recorded_guardrail_information from litellm.proxy.utils import ProxyLogging, _streamable_post_call_pipelines, stream_gated_guardrail_names from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import ContentFilterGuardrail from litellm.types.guardrails import BlockedWord, ContentFilterAction, GuardrailEventHooks @@ -2993,3 +2995,276 @@ async def test_streaming_iterator_hook_pipeline_releases_buffered_content_when_a "RaisingAcloseCallback" in message and "RuntimeError" in message and "cleanup failed" not in message for message in _warnings(caplog) ) + + +# --------------------------------------------------------------------------- +# detect-only post_call pipelines streaming live (streaming_buffer_until_moderated: false) +# --------------------------------------------------------------------------- + + +_DETECT_ONLY_STEP: PipelineStep = PipelineStep(guardrail="gr-post", on_pass="allow", on_fail="next") + + +def _live_scan_guardrail( + seen: Dict[str, Any], + *, + guardrail_name: str = "gr-post", + buffer_until_moderated: bool | None = False, + mask_response_content: bool = False, + rewrite: str | None = None, + raises: Exception | None = None, +) -> CustomGuardrail: + class LiveScanGuardrail(CustomGuardrail): + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): + seen["count"] = seen.get("count", 0) + 1 + seen["texts"] = list(inputs.get("texts") or []) + if raises is not None: + raise raises + if rewrite is not None: + return {**inputs, "texts": [rewrite]} + return inputs + + guardrail = LiveScanGuardrail( + guardrail_name=guardrail_name, + event_hook=GuardrailEventHooks.post_call, + default_on=False, + mask_response_content=mask_response_content, + ) + if buffer_until_moderated is not None: + guardrail.streaming_buffer_until_moderated = buffer_until_moderated + return guardrail + + +def _delta_contents(chunks: List[Any]) -> List[tuple[str | None, str | None]]: + return [(chunk.choices[0].delta.content, chunk.choices[0].finish_reason) for chunk in chunks] + + +def _post_call_statuses(data: Dict[str, Any]) -> List[str]: + return [str(entry.get("guardrail_status")) for entry in recorded_guardrail_information(data)] + + +async def _deliver_with_scan_counts( + proxy_logging: ProxyLogging, user_api_key_dict: UserAPIKeyAuth, data: Dict[str, Any], seen: Dict[str, Any] +) -> tuple[List[Any], List[int]]: + delivered: List[Any] = [] + scans_before_delivery: List[int] = [] + async for item in proxy_logging.async_post_call_streaming_iterator_hook( + user_api_key_dict=user_api_key_dict, response=_async_chunk_iter(_stream_chunks()), request_data=data + ): + delivered.append(item) + scans_before_delivery.append(seen.get("count", 0)) + return delivered, scans_before_delivery + + +@pytest.mark.asyncio +async def test_streaming_iterator_hook_live_pipeline_releases_every_chunk_before_the_end_of_stream_scan( + proxy_logging: ProxyLogging, + make_user_api_key_auth: Callable[..., UserAPIKeyAuth], + monkeypatch: pytest.MonkeyPatch, +) -> None: + seen: Dict[str, Any] = {} + monkeypatch.setattr(litellm, "callbacks", [_live_scan_guardrail(seen)]) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False) + data = _post_call_pipeline_data(step=_DETECT_ONLY_STEP, stream=True) + + delivered, scans_before_delivery = await _deliver_with_scan_counts( + proxy_logging, make_user_api_key_auth(request_route="/v1/chat/completions"), data, seen + ) + + assert scans_before_delivery == [0, 0] + assert _delta_contents(delivered) == [("hello ", None), ("world", "stop")] + assert seen["count"] == 1 + assert seen["texts"] == ["hello world"] + assert _post_call_statuses(data) == ["success"] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("step", "guardrail_kwargs"), + ( + pytest.param(PipelineStep(guardrail="gr-post", on_pass="allow", on_fail="block"), {}, id="on_fail-block"), + pytest.param( + PipelineStep(guardrail="gr-post", on_pass="allow", on_fail="next", on_error="block"), + {}, + id="on_error-block", + ), + pytest.param( + PipelineStep(guardrail="gr-post", on_pass="allow", on_fail="modify_response"), + {}, + id="on_fail-modify_response", + ), + pytest.param(_DETECT_ONLY_STEP, {"buffer_until_moderated": None}, id="flag-unset"), + pytest.param(_DETECT_ONLY_STEP, {"buffer_until_moderated": True}, id="flag-true"), + pytest.param(_DETECT_ONLY_STEP, {"mask_response_content": True}, id="mask-response-content"), + ), +) +async def test_streaming_iterator_hook_pipeline_keeps_buffering_unless_a_detect_only_step_opted_out( + proxy_logging: ProxyLogging, + make_user_api_key_auth: Callable[..., UserAPIKeyAuth], + monkeypatch: pytest.MonkeyPatch, + step: PipelineStep, + guardrail_kwargs: Dict[str, Any], +) -> None: + seen: Dict[str, Any] = {} + monkeypatch.setattr(litellm, "callbacks", [_live_scan_guardrail(seen, **guardrail_kwargs)]) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False) + data = _post_call_pipeline_data(step=step, stream=True) + + delivered, scans_before_delivery = await _deliver_with_scan_counts( + proxy_logging, make_user_api_key_auth(request_route="/v1/chat/completions"), data, seen + ) + + assert scans_before_delivery == [1, 1] + assert _delta_contents(delivered) == [("hello ", None), ("world", "stop")] + + +@pytest.mark.asyncio +async def test_streaming_iterator_hook_live_pipeline_discards_each_rewrite_before_the_next_step_scans( + proxy_logging: ProxyLogging, + make_user_api_key_auth: Callable[..., UserAPIKeyAuth], + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, +) -> None: + first_seen: Final[Dict[str, Any]] = {} + second_seen: Final[Dict[str, Any]] = {} + monkeypatch.setattr( + litellm, + "callbacks", + [ + _live_scan_guardrail(first_seen, guardrail_name="gr-first", rewrite="hello [MASKED]"), + _live_scan_guardrail(second_seen, guardrail_name="gr-second"), + ], + ) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False) + pipeline: Final = GuardrailPipeline( + mode="post_call", + steps=[ + PipelineStep(guardrail="gr-first", on_pass="next", on_fail="next"), + PipelineStep(guardrail="gr-second", on_pass="allow", on_fail="next"), + ], + ) + data: Final = _post_call_pipeline_data(stream=True) + data["metadata"]["_guardrail_pipelines"] = [("response-governance", pipeline)] + data["metadata"]["_pipeline_managed_guardrails"] = {"gr-first", "gr-second"} + + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + delivered, scans_before_delivery = await _deliver_with_scan_counts( + proxy_logging, make_user_api_key_auth(request_route="/v1/chat/completions"), data, second_seen + ) + + assert scans_before_delivery == [0, 0] + assert _delta_contents(delivered) == [("hello ", None), ("world", "stop")] + assert first_seen["texts"] == ["hello world"] + assert second_seen["texts"] == ["hello world"] + assert any( + "'gr-first'" in message and "already received the stream live" in message and "discarded" in message + for message in _warnings(caplog) + ) + + +@pytest.mark.asyncio +async def test_streaming_iterator_hook_live_pipeline_scans_what_a_disconnected_client_received( + proxy_logging: ProxyLogging, + make_user_api_key_auth: Callable[..., UserAPIKeyAuth], + monkeypatch: pytest.MonkeyPatch, +) -> None: + seen: Dict[str, Any] = {} + monkeypatch.setattr(litellm, "callbacks", [_live_scan_guardrail(seen)]) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False) + data = _post_call_pipeline_data(step=_DETECT_ONLY_STEP, stream=True) + stream = 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(_stream_chunks()), + request_data=data, + ) + + first = await stream.__anext__() + await stream.aclose() + + assert first.choices[0].delta.content == "hello " + assert seen["texts"] == ["hello "] + assert _post_call_statuses(data) == ["success"] + + +@pytest.mark.asyncio +async def test_streaming_iterator_hook_live_pipeline_records_a_scan_that_failed_after_the_client_disconnected( + proxy_logging: ProxyLogging, + make_user_api_key_auth: Callable[..., UserAPIKeyAuth], + monkeypatch: pytest.MonkeyPatch, +) -> None: + seen: Dict[str, Any] = {} + monkeypatch.setattr(litellm, "callbacks", [_live_scan_guardrail(seen, raises=RuntimeError("scanner down"))]) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False) + data = _post_call_pipeline_data(step=_DETECT_ONLY_STEP, stream=True) + stream = 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(_stream_chunks()), + request_data=data, + ) + + await stream.__anext__() + await stream.aclose() + + assert seen["count"] == 1 + assert _post_call_statuses(data) == ["guardrail_failed_to_respond"] + + +async def _provider_error_after_first_chunk(chunks: List[Any]) -> AsyncGenerator[Any, None]: + yield chunks[0] + raise RuntimeError("provider dropped the stream") + + +@pytest.mark.asyncio +async def test_streaming_iterator_hook_live_pipeline_scans_what_the_client_received_before_a_provider_error( + proxy_logging: ProxyLogging, + make_user_api_key_auth: Callable[..., UserAPIKeyAuth], + monkeypatch: pytest.MonkeyPatch, +) -> None: + seen: Final[Dict[str, Any]] = {} + monkeypatch.setattr(litellm, "callbacks", [_live_scan_guardrail(seen)]) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False) + data: Final = _post_call_pipeline_data(step=_DETECT_ONLY_STEP, stream=True) + stream: Final = proxy_logging.async_post_call_streaming_iterator_hook( + user_api_key_dict=make_user_api_key_auth(request_route="/v1/chat/completions"), + response=_provider_error_after_first_chunk(_stream_chunks()), + request_data=data, + ) + + first: Final = await stream.__anext__() + with pytest.raises(RuntimeError, match="provider dropped the stream"): + await stream.__anext__() + + assert first.choices[0].delta.content == "hello " + assert seen["texts"] == ["hello "] + assert _post_call_statuses(data) == ["success"] + + +class _EndingRaisesTranslation(OpenAIChatCompletionsHandler): + def released_stream_as_ended(self, responses_so_far): + raise RuntimeError("synthetic translation failure") + + +@pytest.mark.asyncio +async def test_live_pipeline_disconnect_scan_that_cannot_run_records_every_step_guardrail_as_failed( + proxy_logging: ProxyLogging, + make_user_api_key_auth: Callable[..., UserAPIKeyAuth], + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, +) -> None: + seen: Dict[str, Any] = {} + monkeypatch.setattr(litellm, "callbacks", [_live_scan_guardrail(seen)]) + data = _post_call_pipeline_data(step=_DETECT_ONLY_STEP, stream=True) + + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + await proxy_logging._scan_pipeline_stream_cut_short( + released=_stream_chunks()[:1], + user_api_key_dict=make_user_api_key_auth(request_route="/v1/chat/completions"), + request_data=data, + pipelines=tuple(data["metadata"]["_guardrail_pipelines"]), + translation=("completion", _EndingRaisesTranslation()), + ) + + assert seen.get("count") is None + assert _post_call_statuses(data) == ["guardrail_failed_to_respond"] + assert any("cut short" in message and "RuntimeError" in message for message in _warnings(caplog)) +