fix(policy_engine): record a streaming pipeline step once and in the applied guardrails header

CustomGuardrail.__init_subclass__ wrapped _StreamRewriteObserver.apply_guardrail in log_guardrail_information, so every streaming step recorded a second standard_logging_guardrail_information entry and span next to the inner guardrail's own. The observer's method now carries the marker that skips the wrapper. The step also adds the guardrail to the applied guardrails header the way the non-streaming unified path does, so streamed spend rows name the guardrail that scanned them
This commit is contained in:
mateo-berri 2026-09-07 21:48:19 -07:00
parent 69d2ac1edb
commit 91fc1b2010
2 changed files with 65 additions and 6 deletions

View file

@ -7,19 +7,21 @@ pass/fail actions (allow, block, next, modify_response) and data forwarding.
import copy
import time
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Literal
from collections.abc import Callable, Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Literal, TypeVar
from pydantic import BaseModel
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import LOGS_GUARDRAIL_INFORMATION_MARKER
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
ModifyResponseException,
)
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import independent_snapshot
from litellm.proxy.common_utils.callback_utils import add_guardrail_to_applied_guardrails_header
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
UnifiedLLMGuardrails,
)
@ -72,13 +74,23 @@ def _rewrote(sent: tuple[object, ...] | None, returned: tuple[object, ...] | Non
return sent is not None and returned is not None and returned != sent
_GuardrailMethodT = TypeVar("_GuardrailMethodT", bound=Callable[..., object])
def _logged_by_inner_guardrail(method: _GuardrailMethodT) -> _GuardrailMethodT:
vars(method)[LOGS_GUARDRAIL_INFORMATION_MARKER] = True # rebind-ok: stamps the method the class body just defined
return method
class _StreamRewriteObserver(CustomGuardrail):
"""Stand-in handed to the endpoint translation in place of a streaming pipeline step's
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
rewrites are deliverable on translations that write them back across the buffered chunks
(``delivers_ended_stream_text_rewrites``); tool-call rewrites and text rewrites on any
other translation are discarded by the executor, which releases the original chunks."""
other 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``."""
def __init__(self, inner: CustomGuardrail) -> None:
super().__init__(guardrail_name=inner.guardrail_name)
@ -89,6 +101,7 @@ class _StreamRewriteObserver(CustomGuardrail):
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,
@ -305,9 +318,11 @@ class PipelineExecutor:
)
except UndeliverableStreamRewrite:
_release_original_chunks(step.guardrail, streaming_chunks, originals)
return
if observer.rewrote_tool_calls or (observer.rewrote_texts and not deliver_rewrites):
_release_original_chunks(step.guardrail, streaming_chunks, originals)
else:
if observer.rewrote_tool_calls or (observer.rewrote_texts and not deliver_rewrites):
_release_original_chunks(step.guardrail, streaming_chunks, originals)
if not callback.records_own_guardrail_information:
add_guardrail_to_applied_guardrails_header(request_data=hook_input, guardrail_name=step.guardrail)
@staticmethod
async def _run_step(

View file

@ -1099,6 +1099,50 @@ async def test_streaming_step_discards_tool_call_rewrite_and_restores_written_te
assert chunks == [_chunk()]
class _BlockingStreamGuardrail(CustomGuardrail):
def __init__(self):
super().__init__(guardrail_name="masker", event_hook="post_call", default_on=True)
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
raise HTTPException(status_code=400, detail={"error": "output blocked"})
def _recorded_guardrail_statuses(result):
return [
entry["guardrail_status"]
for entry in result.modified_data["metadata"]["standard_logging_guardrail_information"]
]
@pytest.mark.asyncio
async def test_streaming_step_records_guardrail_information_once_on_mask(monkeypatch):
monkeypatch.setattr(litellm, "callbacks", [_TextReturningGuardrail(["hello [MASKED]"])])
result = await _run_streaming_step(_WritingTranslation(), [_chunk()])
assert result.terminal_action == "allow"
assert _recorded_guardrail_statuses(result) == ["success"]
@pytest.mark.asyncio
async def test_streaming_step_records_the_guardrail_in_the_applied_guardrails_header(monkeypatch):
monkeypatch.setattr(litellm, "callbacks", [_TextReturningGuardrail(["hello [MASKED]"])])
result = await _run_streaming_step(_WritingTranslation(), [_chunk()])
assert result.modified_data["metadata"]["applied_guardrails"] == ["masker"]
@pytest.mark.asyncio
async def test_streaming_step_records_guardrail_information_once_on_block(monkeypatch):
monkeypatch.setattr(litellm, "callbacks", [_BlockingStreamGuardrail()])
result = await _run_streaming_step(_WritingTranslation(), [_chunk()])
assert [step.outcome for step in result.step_results] == ["fail"]
assert _recorded_guardrail_statuses(result) == ["guardrail_intervened"]
@pytest.mark.asyncio
async def test_streaming_step_restores_chunks_when_translation_refuses_the_rewrite(monkeypatch, caplog):
monkeypatch.setattr(litellm, "callbacks", [_TextReturningGuardrail(["hello [MASKED]"])])