mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
08b60c409a
commit
c6f5763443
9 changed files with 505 additions and 87 deletions
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue