diff --git a/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py b/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py index 578825d971e..096ec729b6e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py +++ b/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py @@ -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 diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/crowdstrike_aidr.py b/litellm/types/proxy/guardrails/guardrail_hooks/crowdstrike_aidr.py index df1caab6af6..bbd1269bef6 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/crowdstrike_aidr.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/crowdstrike_aidr.py @@ -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, diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py b/tests/unit/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py index c95f7123221..0f3997ef1ee 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py @@ -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