mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
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:
parent
a426dc43cb
commit
c64e746b71
2 changed files with 67 additions and 1 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue