fix(content_filter): log only scan time as streaming post_call guardrail duration (#40760)

The streaming iterator hook timed the whole provider stream and logged that as the
guardrail duration, so PrometheusLogger added LLM generation time to
litellm_overhead_with_guardrails_latency_metric. The hook now accumulates the time
spent inside _filter_single_text per chunk and logs that sum, keeping start_time and
end_time as the wall-clock window.

Resolves LIT-7589

Co-authored-by: yassin <yassin@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-11 14:34:21 -07:00 committed by GitHub
parent a426dc43cb
commit c64e746b71
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 67 additions and 1 deletions

View file

@ -9,6 +9,7 @@ import asyncio
import json
import os
import re
import time
from collections.abc import AsyncGenerator, Coroutine, Mapping, Sequence
from datetime import datetime
from re import Pattern
@ -1702,6 +1703,7 @@ class ContentFilterGuardrail(CustomGuardrail):
start_time: datetime,
masked_entity_count: dict[str, int],
exception_str: str,
duration: float | None = None,
) -> None:
"""
Log guardrail information to request_data metadata.
@ -1713,6 +1715,7 @@ class ContentFilterGuardrail(CustomGuardrail):
start_time: Start time of guardrail execution
masked_entity_count: Count of masked entities by type
exception_str: Exception string if guardrail failed
duration: Seconds spent inside the guardrail; defaults to the wall clock since start_time
"""
# Convert TypedDict detections to regular dicts for JSON serialization
guardrail_json_response: Exception | str | dict | list[dict] = [dict(detection) for detection in detections]
@ -1741,7 +1744,7 @@ class ContentFilterGuardrail(CustomGuardrail):
guardrail_status=status,
start_time=start_time.timestamp(),
end_time=datetime.now().timestamp(),
duration=(datetime.now() - start_time).total_seconds(),
duration=(datetime.now() - start_time).total_seconds() if duration is None else duration,
masked_entity_count=masked_entity_count,
tracing_detail=GuardrailTracingDetail(**tracing_kw),
)
@ -1971,6 +1974,7 @@ class ContentFilterGuardrail(CustomGuardrail):
buffer_size: Final = 50 # Increased buffer to catch patterns split across many chunks
start_time: Final = datetime.now()
scan_seconds: float = 0.0 # rebind-ok: accumulates per-chunk scan time across the stream
detections: list[ContentFilterDetection] = []
masked_entity_count: Final[dict[str, int]] = {}
status: GuardrailStatus = "success"
@ -2007,6 +2011,7 @@ class ContentFilterGuardrail(CustomGuardrail):
# Add a space at the end if it's the final chunk to trigger word boundaries (\b)
text_to_scan = text_to_check + (" " if is_final else "")
choice_detections: list[ContentFilterDetection] = []
scan_started = time.perf_counter()
try:
# _filter_single_text scans the whole accumulated
@ -2024,6 +2029,8 @@ class ContentFilterGuardrail(CustomGuardrail):
except Exception as e:
verbose_proxy_logger.error("ContentFilterGuardrail: Error in masking: %s", e)
masked_text = text_to_scan # Fallback to current text
finally:
scan_seconds += time.perf_counter() - scan_started
# Determine how much can be safely yielded
if is_final:
@ -2074,6 +2081,7 @@ class ContentFilterGuardrail(CustomGuardrail):
start_time=start_time,
masked_entity_count=masked_entity_count,
exception_str=exception_str,
duration=scan_seconds,
)
@staticmethod

View file

@ -626,6 +626,64 @@ class TestContentFilterGuardrail:
assert entry["guardrail_status"] == "success"
assert entry["guardrail_response"] == []
@pytest.mark.asyncio
async def test_streaming_hook_duration_excludes_provider_wait(self):
"""
Streaming post-call: the logged guardrail duration must only cover the
per-chunk scans, not the time spent waiting on the provider between
chunks. PrometheusLogger adds post_call guardrail duration to
litellm_overhead_with_guardrails_latency_metric, so a duration spanning
the whole stream reports LLM generation time as guardrail overhead.
"""
import asyncio
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices
guardrail = ContentFilterGuardrail(
guardrail_name="test-streaming-duration",
patterns=[
ContentFilterPattern(
pattern_type="prebuilt",
pattern_name="email",
action=ContentFilterAction.MASK,
),
],
event_hook=GuardrailEventHooks.post_call,
)
provider_wait_per_chunk = 0.15
chunks = ("Hello ", "world, reach me at ", "test@example.com ")
async def slow_stream():
for i, text in enumerate(chunks):
await asyncio.sleep(provider_wait_per_chunk)
yield ModelResponseStream(
id=f"chunk{i}",
choices=[
StreamingChoices(
delta=Delta(content=text),
index=0,
finish_reason="stop" if i == len(chunks) - 1 else None,
)
],
model="gpt-4",
)
request_data = {"messages": [{"role": "user", "content": "Hi"}], "model": "gpt-4o", "metadata": {}}
async for _ in guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=MagicMock(),
response=slow_stream(),
request_data=request_data,
):
pass
entry = request_data["metadata"]["standard_logging_guardrail_information"][0]
stream_wall_clock = entry["end_time"] - entry["start_time"]
assert stream_wall_clock >= provider_wait_per_chunk * len(chunks)
assert entry["masked_entity_count"].get("email", 0) >= 1
assert 0 < entry["duration"] < provider_wait_per_chunk, entry["duration"]
@pytest.mark.asyncio
async def test_streaming_hook_logs_guardrail_information_mask(self):
"""