mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(guardrails): default CrowdStrike AIDR guardrail to run at end of stream
The CrowdStrike AIDR guardrail has always been intended to work with the complete response, not with streamed chunks. The introduction of `streaming_end_of_stream_only` and the like have caused the guardrail to begin receiving chunked content, which then gets reflected in the CrowdStrike AIDR platform as multiple incomplete findings/events, which looks like a regression to customers. This patch updates the `streaming_end_of_stream_only` default to `True` for this guardrail.
This commit is contained in:
parent
02f61c9c42
commit
599065386d
3 changed files with 29 additions and 22 deletions
|
|
@ -275,7 +275,7 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
|
|||
api_key (str | None): The CrowdStrike AIDR API key. Reads from CS_AIDR_TOKEN env var if None.
|
||||
api_base (str | None): The CrowdStrike AIDR API base URL. Reads from CS_AIDR_BASE_URL env var if None.
|
||||
streaming_end_of_stream_only (bool | None): Scan streamed output once at end of stream instead of
|
||||
every streaming_sampling_rate chunks. Defaults to False.
|
||||
every streaming_sampling_rate chunks. Defaults to True.
|
||||
streaming_sampling_rate (int | None): Scan the accumulated streamed output every Nth chunk. Defaults to 5.
|
||||
async_handler (AsyncHTTPHandler | None): HTTP client to call AI Guard with. Defaults to the shared
|
||||
guardrail-callback client.
|
||||
|
|
@ -316,7 +316,11 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
|
|||
def _set_streaming_params(self, streaming_params: CrowdStrikeAIDRGuardrailConfigModelOptionalParams) -> None:
|
||||
self.streaming_buffer_until_moderated: bool = streaming_params.streaming_buffer_until_moderated or False
|
||||
self.streaming_buffer_release_on_scan: bool = streaming_params.streaming_buffer_release_on_scan or False
|
||||
self.streaming_end_of_stream_only: bool = streaming_params.streaming_end_of_stream_only or False
|
||||
self.streaming_end_of_stream_only: bool = (
|
||||
True
|
||||
if streaming_params.streaming_end_of_stream_only is None
|
||||
else streaming_params.streaming_end_of_stream_only
|
||||
)
|
||||
self.streaming_sampling_rate: int = streaming_params.streaming_sampling_rate or 5
|
||||
|
||||
@override
|
||||
|
|
|
|||
|
|
@ -14,9 +14,9 @@ class CrowdStrikeAIDRGuardrailConfigModelOptionalParams(BaseModel):
|
|||
)
|
||||
streaming_end_of_stream_only: bool | None = Field(
|
||||
default=None,
|
||||
description="If False (default when unset), post_call scans the accumulated streamed response every "
|
||||
"streaming_sampling_rate chunks and an in-flight block stops the stream. If True, the guard runs once "
|
||||
"over the assembled response at end of stream, so flagged content may already have reached the client.",
|
||||
description="If True (default when unset), the guard runs once over the assembled response at end of "
|
||||
"stream, so flagged content may already have reached the client. If False, post_call scans the "
|
||||
"accumulated streamed response every streaming_sampling_rate chunks and an in-flight block stops the stream.",
|
||||
)
|
||||
streaming_sampling_rate: int | None = Field(
|
||||
default=None,
|
||||
|
|
|
|||
|
|
@ -1624,7 +1624,7 @@ def test_initialize_guardrail_defaults_streaming_params() -> None:
|
|||
|
||||
assert handler.streaming_buffer_until_moderated is False
|
||||
assert handler.streaming_buffer_release_on_scan is False
|
||||
assert handler.streaming_end_of_stream_only is False
|
||||
assert handler.streaming_end_of_stream_only is True
|
||||
assert handler.streaming_sampling_rate == 5
|
||||
|
||||
|
||||
|
|
@ -1685,7 +1685,7 @@ def _stream_chunk(content: str, finish_reason: str | None) -> ModelResponseStrea
|
|||
)
|
||||
|
||||
|
||||
async def _guard_calls_for_stream(handler: CrowdStrikeAIDRHandler, chunk_texts: list[str]) -> int:
|
||||
async def _guard_scans_for_stream(handler: CrowdStrikeAIDRHandler, chunk_texts: list[str]) -> list[str]:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import UnifiedLLMGuardrails
|
||||
|
||||
|
|
@ -1693,11 +1693,11 @@ async def _guard_calls_for_stream(handler: CrowdStrikeAIDRHandler, chunk_texts:
|
|||
for i, content in enumerate(chunk_texts):
|
||||
yield _stream_chunk(content, "stop" if i == len(chunk_texts) - 1 else None)
|
||||
|
||||
calls = 0
|
||||
scanned: list[str] = []
|
||||
|
||||
def _allow(request: httpx.Request) -> httpx.Response:
|
||||
nonlocal calls
|
||||
calls += 1
|
||||
messages = json.loads(request.content)["guard_input"]["messages"]
|
||||
scanned.append("".join(message["content"] for message in messages))
|
||||
return httpx.Response(
|
||||
status_code=200, json={"result": {"blocked": False, "transformed": False}}, request=request
|
||||
)
|
||||
|
|
@ -1716,29 +1716,32 @@ async def _guard_calls_for_stream(handler: CrowdStrikeAIDRHandler, chunk_texts:
|
|||
request_data=request_data,
|
||||
):
|
||||
pass
|
||||
return calls
|
||||
return scanned
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("configured", "expected_calls"),
|
||||
("configured", "expected_scans"),
|
||||
[
|
||||
({}, 2),
|
||||
({"streaming_sampling_rate": 2}, 5),
|
||||
({"streaming_end_of_stream_only": True}, 1),
|
||||
({"streaming_end_of_stream_only": True, "streaming_sampling_rate": 2}, 1),
|
||||
({}, ("ABCDEFGHIJ",)),
|
||||
({"streaming_end_of_stream_only": False}, ("ABCDE", "ABCDEFGHIJ")),
|
||||
(
|
||||
{"streaming_end_of_stream_only": False, "streaming_sampling_rate": 2},
|
||||
("AB", "ABCD", "ABCDEF", "ABCDEFGH", "ABCDEFGHIJ"),
|
||||
),
|
||||
({"streaming_end_of_stream_only": True}, ("ABCDEFGHIJ",)),
|
||||
({"streaming_end_of_stream_only": True, "streaming_sampling_rate": 2}, ("ABCDEFGHIJ",)),
|
||||
],
|
||||
)
|
||||
async def test_streaming_params_from_config_control_output_scan_cadence(
|
||||
configured: dict[str, object], expected_calls: int
|
||||
configured: dict[str, object], expected_scans: tuple[str, ...]
|
||||
) -> None:
|
||||
"""10 chunks: default samples at 5 and 10, rate 2 samples 5 times, end-of-stream scans once.
|
||||
|
||||
The final pass is skipped because chunk 10 already scanned the complete output.
|
||||
"""
|
||||
handler = _initialize_from_config(mode="post_call", **configured)
|
||||
|
||||
assert await _guard_calls_for_stream(handler, list("ABCDEFGHIJ")) == expected_calls
|
||||
scanned = await _guard_scans_for_stream(handler, list("ABCDEFGHIJ"))
|
||||
|
||||
assert tuple(scanned) == expected_scans
|
||||
assert scanned[-1] == "ABCDEFGHIJ"
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue