feat(guardrails): run legacy post-call hooks as streaming pipeline steps

A post_call pipeline step whose guardrail only implements the older
async_post_call_success_hook used to skip the stream entirely: PR #38721
fails that shape open with a warning. The streaming step now assembles the
buffered stream into the response the hook expects, runs the hook, ends the
stream with the hook's exception when it raises, and delivers the hook's
rewrite through the same event write-back the unified guardrails use on
chat, Responses, and Messages streams (Messages gets the Anthropic shape).
A stream a pipeline manages no longer runs the same hook again after the
stream ends. A guardrail with neither the unified interface nor a post-call
hook keeps the fail-open, as does a rewrite the buffer cannot be patched
with.
This commit is contained in:
mateo-berri 2026-09-08 12:04:52 -07:00
parent 08b60c409a
commit c6f5763443
9 changed files with 505 additions and 87 deletions

View file

@ -176,6 +176,11 @@ class AnthropicMessagesHandler(BaseTranslation):
super().__init__()
self.adapter = LiteLLMAnthropicMessagesAdapter()
def post_call_hook_response(self, response: object) -> object:
if not isinstance(response, ModelResponse):
return response
return self.adapter.translate_openai_response_to_anthropic(response)
@staticmethod
def _build_streaming_usage_response(
responses_so_far: Sequence[object],

View file

@ -60,6 +60,13 @@ class BaseTranslation(ABC):
text rewrites on every other translation, are undeliverable: the pipeline
executor discards them and releases the original chunks."""
def post_call_hook_response(self, response: object) -> object:
"""The ``response`` this endpoint's non-streaming post-call hooks receive, derived from
the object the translation stores under ``request_data["response"]`` while scanning an
ended stream. Chat and Responses scan that shape already; a translation that scans a
different one (Messages scans an OpenAI-shaped ModelResponse) overrides this."""
return response
@staticmethod
def transform_user_api_key_dict_to_metadata(
user_api_key_dict: Any | None,

View file

@ -3328,9 +3328,10 @@ class ProxyBaseLLMRequestProcessing:
has completed.
Guardrails routed through unified_guardrail are skipped, since they already ran
via its streaming iterator. Guardrails that override
async_post_call_success_hook directly run here, including those that implement
apply_guardrail but keep their native lifecycle hooks.
via its streaming iterator, and so are guardrails a post_call policy pipeline
manages, since the pipeline ran them against the buffered stream. Guardrails
that override async_post_call_success_hook directly run here, including those
that implement apply_guardrail but keep their native lifecycle hooks.
This is audit-only — content has already been delivered to the client.
@ -3340,12 +3341,18 @@ class ProxyBaseLLMRequestProcessing:
_response = assembled_response
try:
from litellm.proxy.proxy_server import llm_router as _global_llm_router
from litellm.proxy.utils import _check_and_merge_model_level_guardrails
from litellm.proxy.utils import (
_check_and_merge_model_level_guardrails,
pipeline_managed_guardrail_names,
)
guardrail_data = _check_and_merge_model_level_guardrails(data=captured_data, llm_router=_global_llm_router)
pipeline_managed: Final = pipeline_managed_guardrail_names(captured_data, "post_call")
for cb in litellm.callbacks:
if not isinstance(cb, CustomGuardrail):
continue
if cb.guardrail_name in pipeline_managed:
continue
if not cb.should_run_guardrail(
data=guardrail_data,
event_type=GuardrailEventHooks.post_call,

View file

@ -121,6 +121,80 @@ class _StreamRewriteObserver(CustomGuardrail):
return outputs
class _ScannedTextRecorder(CustomGuardrail):
def __init__(self, guardrail_name: str) -> None:
super().__init__(guardrail_name=guardrail_name)
self.texts: tuple[str, ...] | None = None
@_logged_by_inner_guardrail
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict, # mutable-ok: matches CustomGuardrail.apply_guardrail
input_type: Literal["request", "response"],
logging_obj: "LiteLLMLoggingObj | None" = None,
) -> GenericGuardrailAPIInputs:
self.texts = _text_snapshot(inputs.get("texts"))
return inputs
class _LegacyHookStreamAdapter(CustomGuardrail):
"""Runs a guardrail that only implements the legacy post-call hook (no unified
``apply_guardrail``, or ``use_native_lifecycle_hooks``) as a streaming pipeline step. The
endpoint translation hands it the texts it scanned plus the assembled response under
``request_data["response"]``; the hook gets that response in the shape its route gives
non-streaming hooks, an exception it raises ends the stream through the executor's
fail/error classification, and a replacement response is re-scanned by the same translation
so its texts reach the client through the translation's ended-stream write-back. A
replacement whose scanned texts do not line up with the originals is undeliverable, so the
executor releases the original chunks."""
def __init__(
self,
inner: CustomGuardrail,
endpoint_translation: "BaseTranslation",
user_api_key_dict: "UserAPIKeyAuth",
) -> None:
super().__init__(guardrail_name=inner.guardrail_name)
self.inner: Final = inner
self.endpoint_translation: Final = endpoint_translation
self.user_api_key_dict: Final = user_api_key_dict
def structured_messages_cover_full_request(self) -> bool:
return self.inner.structured_messages_cover_full_request()
@_logged_by_inner_guardrail
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict, # mutable-ok: matches CustomGuardrail.apply_guardrail
input_type: Literal["request", "response"],
logging_obj: "LiteLLMLoggingObj | None" = None,
) -> GenericGuardrailAPIInputs:
replacement: Final = await self.inner.async_post_call_success_hook(
data=request_data,
user_api_key_dict=self.user_api_key_dict,
response=self.endpoint_translation.post_call_hook_response(request_data.get("response")),
)
if replacement is None:
return inputs
scanned: Final = _text_snapshot(inputs.get("texts"))
rewritten: Final = await self._scanned_texts(replacement, logging_obj)
if scanned is None or rewritten is None or len(rewritten) != len(scanned):
raise UndeliverableStreamRewrite(self.guardrail_name or "unknown")
return {**inputs, "texts": list(rewritten)}
async def _scanned_texts(self, response: object, logging_obj: "LiteLLMLoggingObj | None") -> tuple[str, ...] | None:
recorder: Final = _ScannedTextRecorder(self.guardrail_name or "unknown")
await self.endpoint_translation.process_output_response(
response=response,
guardrail_to_apply=recorder,
litellm_logging_obj=logging_obj,
user_api_key_dict=self.user_api_key_dict,
)
return recorder.texts
def _prepare_hook_input(
step: PipelineStep,
callback: CustomGuardrail,
@ -286,16 +360,23 @@ class PipelineExecutor:
endpoint_translation: "BaseTranslation",
streaming_chunks: list[object], # mutable-ok: shared buffered-stream chunks the translation rewrites in place
hook_input: dict[str, object], # mutable-ok: same request-payload shape as data
user_api_key_dict: "UserAPIKeyAuth | None",
user_api_key_dict: "UserAPIKeyAuth",
litellm_logging_obj: "LiteLLMLoggingObj | None",
) -> None:
"""Run one streaming post_call step through the endpoint translation, delivering
text rewrites on translations that support ended-stream write-back. A rewrite that
cannot reach the client yet (a tool-call rewrite, a text rewrite on a translation
without write-back, or one the translation 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."""
observer: Final = _StreamRewriteObserver(callback)
text 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 (a tool-call rewrite, a text rewrite on a translation without write-back, 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."""
scanner: Final = (
callback
if PipelineExecutor.supports_unified_execution(callback)
else _LegacyHookStreamAdapter(callback, endpoint_translation, user_api_key_dict)
)
observer: Final = _StreamRewriteObserver(scanner)
deliver_rewrites: Final = type(endpoint_translation).delivers_ended_stream_text_rewrites
originals: Final = copy.deepcopy(streaming_chunks)
try:
@ -379,11 +460,11 @@ class PipelineExecutor:
if isinstance(response, dict):
callback.mark_pre_call_hook_ran(response)
elif mode == "post_call" and streaming_chunks is not None:
if not use_unified or endpoint_translation is None:
if endpoint_translation is None:
return (
"error",
None,
f"Guardrail '{step.guardrail}' does not support streaming pipeline execution",
f"Guardrail '{step.guardrail}' cannot run on a stream without an endpoint translation",
None,
)
await PipelineExecutor._run_streaming_step(
@ -433,10 +514,20 @@ class PipelineExecutor:
@staticmethod
def supports_unified_execution(callback: CustomGuardrail) -> bool:
"""Whether this guardrail runs through the unified apply_guardrail path,
the interface streaming pipeline execution requires."""
"""Whether this guardrail runs through the unified apply_guardrail path."""
return "apply_guardrail" in type(callback).__dict__ and not callback.use_native_lifecycle_hooks
@staticmethod
def supports_streaming_execution(callback: CustomGuardrail) -> bool:
"""Whether a streaming pipeline step can run this guardrail against the buffered
stream: through the unified path, or through its own post-call hook on the
assembled response. A guardrail with neither (one that only rewrites the stream
through its iterator hook) has to keep running on its own."""
return (
PipelineExecutor.supports_unified_execution(callback)
or type(callback).async_post_call_success_hook is not CustomLogger.async_post_call_success_hook
)
@staticmethod
def find_guardrail_callback(guardrail_name: str) -> CustomGuardrail | None:
"""Look up an initialized guardrail callback by name from litellm.callbacks."""

View file

@ -451,7 +451,7 @@ def _pipeline_step_guardrail_names(pipelines: Sequence[tuple[str, "GuardrailPipe
return frozenset(step.guardrail for _policy_name, pipeline in pipelines for step in pipeline.steps)
def _pipeline_managed_guardrail_names(
def pipeline_managed_guardrail_names(
data: Mapping[str, object], mode: Literal["pre_call", "post_call"]
) -> frozenset[str]:
return _pipeline_step_guardrail_names(
@ -514,9 +514,9 @@ def _merge_pipeline_metadata_writes(
_merge_pipeline_metadata_bucket(data, bucket_key, modified_data.get(bucket_key))
def _pipeline_step_supports_unified_streaming(guardrail_name: str) -> bool:
def _pipeline_step_supports_streaming(guardrail_name: str) -> bool:
callback: Final = PipelineExecutor.find_guardrail_callback(guardrail_name)
return callback is not None and PipelineExecutor.supports_unified_execution(callback)
return callback is not None and PipelineExecutor.supports_streaming_execution(callback)
def _post_call_pipelines(data: Mapping[str, object]) -> tuple[tuple[str, "GuardrailPipeline"], ...]:
@ -541,14 +541,15 @@ def _warn_background_skips_post_call_pipelines(data: Mapping[str, object]) -> No
def _pipeline_is_streamable(policy_name: str, pipeline: "GuardrailPipeline") -> bool:
unsupported: Final = tuple(
dict.fromkeys(
step.guardrail for step in pipeline.steps if not _pipeline_step_supports_unified_streaming(step.guardrail)
step.guardrail for step in pipeline.steps if not _pipeline_step_supports_streaming(step.guardrail)
)
)
if not unsupported:
return True
verbose_proxy_logger.warning(
"Policy '%s' has post_call pipeline guardrails without the unified apply_guardrail interface, "
"which streaming pipelines need; the stream skips the pipeline and its guardrails run on their own: %s",
"Policy '%s' has post_call pipeline guardrails with neither the unified apply_guardrail interface nor a "
"post-call hook, one of which streaming pipelines need; the stream skips the pipeline and its guardrails "
"run on their own: %s",
policy_name,
", ".join(unsupported),
)
@ -562,11 +563,12 @@ def _streamable_post_call_pipelines(
The post_call pipelines a streaming response can be gated through.
Streaming pipelines scan the buffered stream through the endpoint guardrail
translation of the request route, so every step's guardrail needs the
unified apply_guardrail interface and the route needs a translation. A
pipeline that cannot be run that way yet is left out and its guardrails
run on the stream on their own, the way they did before pipelines ran on
streams at all, with a warning naming the pipeline.
translation of the request route, so every step's guardrail needs either the
unified apply_guardrail interface or a post-call hook to run against the
assembled response, and the route needs a translation. A pipeline that
cannot be run that way yet is left out and its guardrails run on the stream
on their own, the way they did before pipelines ran on streams at all, with
a warning naming the pipeline.
"""
post_call_pipelines: Final = _post_call_pipelines(request_data)
if not post_call_pipelines:
@ -1968,7 +1970,7 @@ class ProxyLogging:
)
# Get pipeline-managed guardrails to skip in normal loop
pipeline_managed: Final = _pipeline_managed_guardrail_names(data, "pre_call")
pipeline_managed: Final = pipeline_managed_guardrail_names(data, "pre_call")
caps: Final = ProxyLogging._callback_capabilities()
# Skip the per-request callback walk entirely when nothing in
@ -2956,7 +2958,7 @@ class ProxyLogging:
if pipeline_response is not None:
response = pipeline_response # rebind-ok: adopt the pipeline's replacement response, same contract as the callback loops below
pipeline_managed: Final = _pipeline_managed_guardrail_names(data, "post_call")
pipeline_managed: Final = pipeline_managed_guardrail_names(data, "post_call")
guardrail_callbacks, other_callbacks = _partition_post_call_callbacks()
try:
# Merge model-level guardrails before checking which guardrails to run
@ -3272,7 +3274,7 @@ class ProxyLogging:
_cached_guardrail_data: dict | None = None
_guardrail_data_computed = False
pipeline_managed: Final = (
_pipeline_managed_guardrail_names(data, "post_call") if caps.has_guardrail else frozenset()
pipeline_managed_guardrail_names(data, "post_call") if caps.has_guardrail else frozenset()
)
for callback in litellm.callbacks:

View file

@ -2156,3 +2156,29 @@ class TestAnthropicMessagesHandlerStreamingScanKey:
assert open_key == StreamingScanKey(texts=("hi",))
assert len(ended_key.tool_calls) == 1 and "get_weather" in ended_key.tool_calls[0]
assert ended_key != open_key
class TestAnthropicMessagesHandlerPostCallHookResponse:
def test_openai_shaped_stream_assembly_reaches_the_hook_as_a_messages_response(self):
from litellm.types.utils import Choices, Message, ModelResponse, Usage
assembled = ModelResponse(
id="msg_1",
model="claude",
choices=[Choices(message=Message(role="assistant", content="hello world"), finish_reason="stop")],
usage=Usage(prompt_tokens=1, completion_tokens=2, total_tokens=3),
)
hook_response = AnthropicMessagesHandler().post_call_hook_response(assembled)
assert hook_response["type"] == "message"
assert hook_response["role"] == "assistant"
assert hook_response["content"] == [{"type": "text", "text": "hello world"}]
assert hook_response["stop_reason"] == "end_turn"
assert hook_response["usage"]["input_tokens"] == 1
assert hook_response["usage"]["output_tokens"] == 2
def test_anything_else_reaches_the_hook_untouched(self):
native = {"type": "message", "role": "assistant", "content": [{"type": "text", "text": "hi"}]}
assert AnthropicMessagesHandler().post_call_hook_response(native) is native

View file

@ -534,7 +534,6 @@ async def test_guardrail_not_found_uses_on_fail(monkeypatch):
],
)
monkeypatch.setattr(litellm, "callbacks", [])
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
@ -1153,3 +1152,152 @@ async def test_streaming_step_restores_chunks_when_translation_refuses_the_rewri
_assert_passed_with_discard_warning(result, caplog)
assert chunks == [_chunk()]
class _LegacyHookGuardrail(CustomGuardrail):
"""A guardrail with only the legacy post-call hook: it never defines apply_guardrail."""
def __init__(self, replacement=None, raises=None):
super().__init__(guardrail_name="masker", event_hook="post_call", default_on=True)
self.replacement = replacement
self.raises = raises
self.calls = []
async def async_post_call_success_hook(self, data, user_api_key_dict, response):
self.calls.append({"data": data, "user_api_key_dict": user_api_key_dict, "response": response})
if self.raises is not None:
raise self.raises
return self.replacement
class _NativeHooksGuardrail(_LegacyHookGuardrail):
use_native_lifecycle_hooks = True
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
raise AssertionError("a guardrail that keeps its native hooks never runs apply_guardrail")
class _LegacyScanningTranslation:
"""Stores the assembled response under request_data["response"] before scanning, like the
chat, Responses, and Messages handlers, hands hooks a route-native shape, and re-extracts one
text per entry of a replacement's "texts"."""
delivers_ended_stream_text_rewrites = True
def post_call_hook_response(self, response):
return {"native": True, "text": response["text"]}
async def process_output_streaming_response(
self,
responses_so_far,
guardrail_to_apply,
litellm_logging_obj=None,
user_api_key_dict=None,
request_data=None,
deliver_ended_stream_rewrites=False,
):
request_data.setdefault("response", {"text": responses_so_far[0]["text"]})
outputs = await guardrail_to_apply.apply_guardrail(
inputs={"texts": [responses_so_far[0]["text"]]},
request_data=request_data,
input_type="response",
logging_obj=litellm_logging_obj,
)
responses_so_far[0]["text"] = outputs["texts"][0]
return responses_so_far
async def process_output_response(
self, response, guardrail_to_apply, litellm_logging_obj=None, user_api_key_dict=None, request_data=None
):
await guardrail_to_apply.apply_guardrail(
inputs={"texts": list(response["texts"])},
request_data={"response": response},
input_type="response",
logging_obj=litellm_logging_obj,
)
return response
async def _run_legacy_streaming_step(monkeypatch, guardrail, chunks, on_fail="block", on_error="next"):
monkeypatch.setattr(litellm, "callbacks", [guardrail])
return await PipelineExecutor.execute_steps(
steps=[PipelineStep(guardrail="masker", on_pass="allow", on_fail=on_fail, on_error=on_error)],
mode="post_call",
data={"model": "m"},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="p",
streaming_chunks=chunks,
endpoint_translation=_LegacyScanningTranslation(),
)
@pytest.mark.asyncio
@pytest.mark.parametrize("guardrail_class", [_LegacyHookGuardrail, _NativeHooksGuardrail])
async def test_streaming_step_runs_legacy_hook_and_delivers_its_rewrite(monkeypatch, caplog, guardrail_class):
guardrail = guardrail_class(replacement={"texts": ["[REWRITTEN] hello world"]})
chunks = [_chunk()]
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
result = await _run_legacy_streaming_step(monkeypatch, guardrail, chunks)
assert result.terminal_action == "allow"
assert [step.outcome for step in result.step_results] == ["pass"]
assert chunks[0]["text"] == "[REWRITTEN] hello world"
assert [call["response"] for call in guardrail.calls] == [{"native": True, "text": "hello world"}]
assert guardrail.calls[0]["data"]["model"] == "m"
assert result.modified_data["metadata"]["applied_guardrails"] == ["masker"]
assert not any("discarded" in record.getMessage() for record in caplog.records)
@pytest.mark.asyncio
async def test_streaming_step_passes_untouched_when_legacy_hook_returns_none(monkeypatch, caplog):
guardrail = _LegacyHookGuardrail(replacement=None)
chunks = [_chunk()]
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
result = await _run_legacy_streaming_step(monkeypatch, guardrail, chunks)
assert result.terminal_action == "allow"
assert len(guardrail.calls) == 1
assert chunks == [_chunk()]
assert not any("discarded" in record.getMessage() for record in caplog.records)
@pytest.mark.skipif(HTTPException is None, reason="fastapi not installed")
@pytest.mark.asyncio
async def test_streaming_step_blocks_with_the_legacy_hook_exception(monkeypatch):
exc = HTTPException(status_code=400, detail={"error": "output blocked"})
chunks = [_chunk()]
result = await _run_legacy_streaming_step(monkeypatch, _LegacyHookGuardrail(raises=exc), chunks)
assert result.terminal_action == "block"
assert [step.outcome for step in result.step_results] == ["fail"]
assert result.original_exception is exc
assert chunks == [_chunk()]
@pytest.mark.asyncio
async def test_streaming_step_takes_on_error_when_legacy_hook_crashes(monkeypatch):
chunks = [_chunk()]
result = await _run_legacy_streaming_step(
monkeypatch, _LegacyHookGuardrail(raises=ValueError("boom")), chunks, on_error="block"
)
assert result.terminal_action == "block"
assert [step.outcome for step in result.step_results] == ["error"]
assert result.step_results[0].error_detail == "boom"
@pytest.mark.asyncio
async def test_streaming_step_discards_legacy_rewrite_whose_texts_do_not_line_up(monkeypatch, caplog):
guardrail = _LegacyHookGuardrail(replacement={"texts": ["split", "in two"]})
chunks = [_chunk()]
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
result = await _run_legacy_streaming_step(monkeypatch, guardrail, chunks)
_assert_passed_with_discard_warning(result, caplog)
assert chunks == [_chunk()]

View file

@ -671,6 +671,32 @@ async def test_deferred_stream_guardrails_run_native_hook_when_opted_out(monkeyp
assert routed.native_hooks_ran == []
@pytest.mark.asyncio
async def test_deferred_stream_guardrails_skip_pipeline_managed_native_hook(monkeypatch):
"""A post_call pipeline step already ran the opted-out guardrail's own hook against
the buffered stream, so the deferred audit must not run it a second time."""
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.types.proxy.policy_engine.pipeline_types import GuardrailPipeline, PipelineStep
from litellm.types.utils import Choices, Message, ModelResponse
pipeline_managed = _KeepsNativeHooks(event_hook=GuardrailEventHooks.post_call, default_on=True)
monkeypatch.setattr(litellm, "callbacks", [pipeline_managed])
pipeline = GuardrailPipeline(mode="post_call", steps=[PipelineStep(guardrail="keeps_native", on_fail="block")])
await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails(
captured_data={
"messages": [{"role": "user", "content": "hi"}],
"metadata": {"_guardrail_pipelines": [("response-governance", pipeline)]},
},
captured_user_api_key_dict=UserAPIKeyAuth(api_key="sk-1234"),
captured_logging_obj=_streaming_logging_obj(),
assembled_response=ModelResponse(choices=[Choices(message=Message(role="assistant", content="hello"))]),
cache_hit=False,
)
assert pipeline_managed.native_hooks_ran == []
@pytest.mark.asyncio
async def test_realtime_guardrails_skip_opted_out_guardrail(monkeypatch):
"""The realtime path calls apply_guardrail directly, so the opt-out has to be

View file

@ -11,6 +11,7 @@ from __future__ import annotations
import asyncio
import json
from copy import deepcopy
import logging
from typing import Any, Callable, Dict, List
from unittest.mock import AsyncMock, MagicMock, patch
@ -1497,29 +1498,78 @@ async def _async_chunk_iter(chunks: List[Any]):
yield chunk
def test_streamable_post_call_pipelines_keeps_supported_and_drops_unsupported(
def _legacy_hook_stream_guardrail(
seen: Dict[str, Any],
rewrite: Callable[[Any], Any] | None = None,
raises: Exception | None = None,
native_lifecycle: bool = False,
) -> CustomGuardrail:
class LegacyHookGuardrail(CustomGuardrail):
use_native_lifecycle_hooks = native_lifecycle
async def async_post_call_success_hook(self, data, user_api_key_dict, response):
seen["count"] = seen.get("count", 0) + 1
seen["data"] = data
seen["user_api_key_dict"] = user_api_key_dict
seen["response"] = deepcopy(response)
if raises is not None:
raise raises
return None if rewrite is None else rewrite(response)
if native_lifecycle:
class NativeLifecycleGuardrail(LegacyHookGuardrail):
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
raise AssertionError("a guardrail that keeps its native hooks never runs apply_guardrail")
return NativeLifecycleGuardrail(
guardrail_name="gr-post", event_hook=GuardrailEventHooks.post_call, default_on=False
)
return LegacyHookGuardrail(guardrail_name="gr-post", event_hook=GuardrailEventHooks.post_call, default_on=False)
def _iterator_hook_only_guardrail(name: str, seen: Dict[str, Any]) -> CustomGuardrail:
class IteratorHookGuardrail(CustomGuardrail):
async def async_post_call_streaming_iterator_hook(self, user_api_key_dict, response, request_data):
seen["count"] = seen.get("count", 0) + 1
async for item in response:
item.choices[0].delta.content = f"[governed] {item.choices[0].delta.content}"
yield item
return IteratorHookGuardrail(guardrail_name=name, event_hook=GuardrailEventHooks.post_call, default_on=True)
def _rewritten_model_response(response: Any) -> litellm.ModelResponse:
payload = response.model_dump()
payload["choices"][0]["message"]["content"] = "[REWRITTEN] " + payload["choices"][0]["message"]["content"]
return litellm.ModelResponse(**payload)
def test_streamable_post_call_pipelines_keeps_hook_guardrails_and_drops_iterator_only(
make_user_api_key_auth, monkeypatch, caplog
):
class NativeOnlyGuardrail(CustomGuardrail):
pass
supported = _unified_stream_guardrail({})
native_only = NativeOnlyGuardrail(guardrail_name="gr-native", event_hook=GuardrailEventHooks.post_call)
monkeypatch.setattr(litellm, "callbacks", [supported, native_only])
governed = GuardrailPipeline(mode="post_call", steps=[PipelineStep(guardrail="gr-post", on_fail="block")])
legacy = _legacy_hook_stream_guardrail({})
legacy.guardrail_name = "gr-legacy"
iterator_only = _iterator_hook_only_guardrail("gr-iterator", {})
monkeypatch.setattr(litellm, "callbacks", [supported, legacy, iterator_only])
governed = GuardrailPipeline(
mode="post_call",
steps=[PipelineStep(guardrail="gr-post", on_fail="next"), PipelineStep(guardrail="gr-legacy", on_fail="block")],
)
ungoverned = GuardrailPipeline(
mode="post_call",
steps=[PipelineStep(guardrail="gr-post", on_fail="next"), PipelineStep(guardrail="gr-native", on_fail="block")],
steps=[PipelineStep(guardrail="gr-post", on_fail="next"), PipelineStep(guardrail="gr-iterator", on_fail="block")],
)
pre_call = GuardrailPipeline(mode="pre_call", steps=[PipelineStep(guardrail="gr-native", on_fail="block")])
pre_call = GuardrailPipeline(mode="pre_call", steps=[PipelineStep(guardrail="gr-iterator", on_fail="block")])
data = {"metadata": {"_guardrail_pipelines": [("governed", governed), ("ungoverned", ungoverned), ("req", pre_call)]}}
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
streamable = _streamable_post_call_pipelines(data, make_user_api_key_auth(request_route="/v1/chat/completions"))
assert streamable == (("governed", governed),)
assert any("'ungoverned'" in message and "gr-native" in message for message in _warnings(caplog))
assert not any("'governed'" in message for message in _warnings(caplog))
assert any("'ungoverned'" in message and "gr-iterator" in message for message in _warnings(caplog))
assert not any("'governed'" in message or "gr-legacy" in message for message in _warnings(caplog))
def test_streamable_post_call_pipelines_is_empty_on_route_without_translation(
@ -1569,55 +1619,123 @@ async def test_pre_call_hook_allows_streaming_when_pipeline_guardrail_supports_u
@pytest.mark.asyncio
@pytest.mark.parametrize("native_lifecycle", [False, True])
async def test_streaming_iterator_hook_releases_stream_when_pipeline_guardrail_lacks_unified_support(
async def test_streaming_iterator_hook_runs_legacy_hook_and_delivers_its_rewrite(
proxy_logging, make_user_api_key_auth, monkeypatch, native_lifecycle, caplog
):
seen: Dict[str, Any] = {}
if native_lifecycle:
class NativeOnlyGuardrail(CustomGuardrail):
use_native_lifecycle_hooks = True
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
seen["count"] = seen.get("count", 0) + 1
return inputs
else:
class NativeOnlyGuardrail(CustomGuardrail):
async def async_post_call_success_hook(self, data, user_api_key_dict, response):
seen["count"] = seen.get("count", 0) + 1
return response
monkeypatch.setattr(
litellm,
"callbacks",
[NativeOnlyGuardrail(guardrail_name="gr-post", event_hook=GuardrailEventHooks.post_call, default_on=False)],
)
guardrail = _legacy_hook_stream_guardrail(seen, rewrite=_rewritten_model_response, native_lifecycle=native_lifecycle)
monkeypatch.setattr(litellm, "callbacks", [guardrail])
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False)
data = _post_call_pipeline_data(stream=True)
chunks = _stream_chunks()
delivered: List[Any] = []
auth = make_user_api_key_auth(request_route="/v1/chat/completions")
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
out = await proxy_logging.pre_call_hook(
user_api_key_dict=make_user_api_key_auth(),
data=data,
call_type="completion",
guardrails_only=True,
user_api_key_dict=auth, data=data, call_type="completion", guardrails_only=True
)
delivered = [
item
async for item in proxy_logging.async_post_call_streaming_iterator_hook(
user_api_key_dict=auth, response=_async_chunk_iter(chunks), request_data=data
)
]
assert out is not None and out.get("stream") is True
assert seen["count"] == 1
assert isinstance(seen["response"], litellm.ModelResponse)
assert seen["response"].choices[0].message.content == "hello world"
assert seen["data"]["messages"] == data["messages"]
assert seen["user_api_key_dict"] is auth
assert [id(item) for item in delivered] == [id(chunk) for chunk in chunks]
assert delivered[0].choices[0].delta.content == "[REWRITTEN] hello world"
assert delivered[1].choices[0].delta.content in (None, "")
assert delivered[1].choices[0].finish_reason == "stop"
assert data["metadata"]["applied_guardrails"] == ["gr-post"]
assert _warnings(caplog) == []
@pytest.mark.asyncio
async def test_streaming_iterator_hook_releases_stream_untouched_when_legacy_hook_returns_none(
proxy_logging, make_user_api_key_auth, monkeypatch
):
seen: Dict[str, Any] = {}
monkeypatch.setattr(litellm, "callbacks", [_legacy_hook_stream_guardrail(seen)])
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False)
data = _post_call_pipeline_data(stream=True)
chunks = _stream_chunks()
delivered = [
item
async for item in proxy_logging.async_post_call_streaming_iterator_hook(
user_api_key_dict=make_user_api_key_auth(request_route="/v1/chat/completions"),
response=_async_chunk_iter(chunks),
request_data=data,
)
]
assert seen["count"] == 1
assert [id(item) for item in delivered] == [id(chunk) for chunk in chunks]
assert [item.choices[0].delta.content for item in delivered] == ["hello ", "world"]
@pytest.mark.asyncio
async def test_streaming_iterator_hook_ends_stream_with_legacy_hook_exception(
proxy_logging, make_user_api_key_auth, monkeypatch
):
seen: Dict[str, Any] = {}
blocked = HTTPException(status_code=400, detail={"error": "output blocked"})
monkeypatch.setattr(litellm, "callbacks", [_legacy_hook_stream_guardrail(seen, raises=blocked)])
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False)
data = _post_call_pipeline_data(stream=True)
delivered: List[Any] = []
async def _drain() -> None:
async for item in proxy_logging.async_post_call_streaming_iterator_hook(
user_api_key_dict=make_user_api_key_auth(request_route="/v1/chat/completions"),
response=_async_chunk_iter(_stream_chunks()),
request_data=data,
):
delivered.append(item)
assert out is not None
assert out.get("stream") is True
assert [item is chunk for item, chunk in zip(delivered, chunks)] == [True, True]
assert len(delivered) == 2
assert seen.get("count") is None
assert any("'response-governance'" in message and "gr-post" in message for message in _warnings(caplog))
with pytest.raises(HTTPException) as info:
await _drain()
assert seen["count"] == 1
assert delivered == []
assert info.value is blocked
@pytest.mark.asyncio
async def test_streaming_iterator_hook_delivers_legacy_hook_rewrite_on_anthropic_sse(
proxy_logging, make_user_api_key_auth, monkeypatch
):
seen: Dict[str, Any] = {}
def rewrite(response: Any) -> Dict[str, Any]:
return {**response, "content": [{"type": "text", "text": "[REWRITTEN] " + response["content"][0]["text"]}]}
monkeypatch.setattr(litellm, "callbacks", [_legacy_hook_stream_guardrail(seen, rewrite=rewrite)])
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False)
data = _post_call_pipeline_data(stream=True)
delivered = [
item
async for item in proxy_logging.async_post_call_streaming_iterator_hook(
user_api_key_dict=make_user_api_key_auth(request_route="/v1/messages"),
response=_async_chunk_iter(_anthropic_sse_chunks()),
request_data=data,
)
]
assert seen["count"] == 1
assert seen["response"]["content"][0]["text"] == "hello world"
assert seen["response"]["role"] == "assistant"
raw = b"".join(delivered).decode()
assert "[REWRITTEN] hello world" in raw
assert raw.count("event: content_block_delta") == 1
for expected_event in ("message_start", "content_block_start", "content_block_stop", "message_delta", "message_stop"):
assert f"event: {expected_event}" in raw
@pytest.mark.asyncio
@ -1625,19 +1743,7 @@ async def test_streaming_iterator_hook_runs_iterator_hook_guardrail_whose_pipeli
proxy_logging, make_user_api_key_auth, monkeypatch, caplog
):
seen: Dict[str, Any] = {}
class IteratorHookGuardrail(CustomGuardrail):
async def async_post_call_streaming_iterator_hook(self, user_api_key_dict, response, request_data):
seen["count"] = seen.get("count", 0) + 1
async for item in response:
item.choices[0].delta.content = f"[governed] {item.choices[0].delta.content}"
yield item
monkeypatch.setattr(
litellm,
"callbacks",
[IteratorHookGuardrail(guardrail_name="gr-post", event_hook=GuardrailEventHooks.post_call, default_on=True)],
)
monkeypatch.setattr(litellm, "callbacks", [_iterator_hook_only_guardrail("gr-post", seen)])
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False)
data = _post_call_pipeline_data(stream=True)