fix(guardrails): unwrap HiddenParamsAsyncIteratorWrapper before deferred dispatch class sniffing

This commit is contained in:
mateo-berri 2026-08-29 02:30:34 -07:00
parent 664697133b
commit 229970c500
2 changed files with 43 additions and 2 deletions

View file

@ -3045,10 +3045,17 @@ class ProxyBaseLLMRequestProcessing:
Raw async generators from passthrough routes bypass all three and
would orphan the closure, so they are not armed here.
The router wraps iterators that cannot carry _hidden_params in
HiddenParamsAsyncIteratorWrapper, so class sniffing runs on the
unwrapped inner iterator.
"""
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
from litellm.router_utils.add_retry_fallback_headers import HiddenParamsAsyncIteratorWrapper
if isinstance(response, CustomStreamWrapper):
unwrapped: Final = response._inner if isinstance(response, HiddenParamsAsyncIteratorWrapper) else response
if isinstance(unwrapped, CustomStreamWrapper):
# Intentionally a live reference (not a copy) — mirrors
# ProxyLogging.post_call_success_hook which also mutates
# data["guardrail_to_apply"] during iteration.
@ -3075,7 +3082,7 @@ class ProxyBaseLLMRequestProcessing:
LiteLLMCompletionStreamingIterator,
)
if isinstance(response, LiteLLMCompletionStreamingIterator):
if isinstance(unwrapped, LiteLLMCompletionStreamingIterator):
_captured_bridge_logging_obj: Final = logging_obj
async def _on_deferred_bridged_stream_complete(assembled_response: object, cache_hit: object) -> None:

View file

@ -1352,6 +1352,40 @@ class TestArmDeferredStreamDispatch:
assert recorded["cache_hit"] is False
assert recorded["prefer_async_handlers"] is True
@pytest.mark.asyncio
async def test_router_wrapped_bridged_iterator_gets_csw_arg_shape(self):
"""The router wraps iterators without _hidden_params in
HiddenParamsAsyncIteratorWrapper before the proxy arms deferral, so
every production streamed /v1/responses reaches arming wrapped;
sniffing the wrapper instead of the inner iterator armed the 1-arg
native closure against the CSW's 2-arg stored shape and leaked a
TypeError 500 frame into the stream."""
from litellm.responses.litellm_completion_transformation.streaming_iterator import (
LiteLLMCompletionStreamingIterator,
)
from litellm.router_utils.add_retry_fallback_headers import (
HiddenParamsAsyncIteratorWrapper,
)
logging_obj, recorded = self._dispatch_recording_logging_obj()
wrapped = HiddenParamsAsyncIteratorWrapper(object.__new__(LiteLLMCompletionStreamingIterator))
self._processor()._arm_deferred_stream_dispatch(
response=wrapped,
route_type="aresponses",
user_api_key_dict=MagicMock(),
logging_obj=logging_obj,
)
assembled = object()
logging_obj._deferred_stream_complete_args = (assembled, False)
ProxyLogging._fire_deferred_stream_logging({"litellm_logging_obj": logging_obj})
await asyncio.sleep(0)
assert recorded["result"] is assembled
assert recorded["cache_hit"] is False
assert recorded["prefer_async_handlers"] is True
@pytest.mark.asyncio
async def test_native_stream_closure_enqueues_single_coroutine(self):
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER