mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(guardrails): unwrap HiddenParamsAsyncIteratorWrapper before deferred dispatch class sniffing
This commit is contained in:
parent
664697133b
commit
229970c500
2 changed files with 43 additions and 2 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue