diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrail/test_unified_guardrail.py b/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrail/test_unified_guardrail.py index f2247b3d032..dfb49a48a06 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrail/test_unified_guardrail.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrail/test_unified_guardrail.py @@ -53,7 +53,7 @@ async def _upstream( def _guarded_stream( - observer: _Observer, upstream: AsyncIterator[ModelResponseStream], *, route: str | None = "/chat/completions" + observer: _Observer, upstream: AsyncIterator[object], *, route: str | None = "/chat/completions" ) -> AsyncIterator[object]: return UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook( user_api_key_dict=UserAPIKeyAuth(api_key="test", request_route=route), @@ -96,6 +96,19 @@ async def test_a_failing_observer_never_breaks_the_stream(warning_messages: Call ] +async def test_a_stream_no_handler_understands_is_forwarded_unobserved( + warning_messages: Callable[[str], list[str]], +) -> None: + observer: Final = _Observer() + + async def raw_bytes() -> AsyncIterator[bytes]: + yield b"data: raw\n\n" + + sent: Final = [chunk async for chunk in _guarded_stream(observer, raw_bytes(), route=None)] + + assert (sent, observer.observed, warning_messages("observe-only stream check")) == ([b"data: raw\n\n"], [], []) + + async def test_an_empty_stream_is_not_observed(warning_messages: Callable[[str], list[str]]) -> None: observer: Final = _Observer()