mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
06006932c7
commit
b1d0e45f8d
12 changed files with 724 additions and 35 deletions
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()))
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue