feat(policy-engine): stream detect-only post_call pipeline steps live when buffering is off (#45516)

* feat(policy-engine): stream detect-only post_call pipeline steps live when buffering is off

* fix(policy-engine): scan live pipeline streams a provider error cuts short and discard rewrites per step

A live detect-only pipeline now scans the chunks the client received when the provider fails mid-stream, then re-raises the error. Each step's rewrite is discarded right after the step, so later steps scan the text the client actually received.

* fix(proxy): keep a guardrail verdict recorded after a stream failure on the failure spend row

---------

Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-09 15:46:51 -07:00 • committed by GitHub
parent 06006932c7
commit b1d0e45f8d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
12 changed files with 724 additions and 35 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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