mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(guardrails): match deferred stream dispatch shape per stream owner and defer passthrough logging until guardrail eos
This commit is contained in:
parent
64eec53fd8
commit
f60ccf6234
6 changed files with 381 additions and 79 deletions
|
|
@ -2341,53 +2341,14 @@ class ProxyBaseLLMRequestProcessing:
|
|||
if requested_model_from_client:
|
||||
self.data["_litellm_client_requested_model"] = requested_model_from_client
|
||||
|
||||
# Streaming: attach a closure that fires after all guardrail
|
||||
# end-of-stream blocks complete. CSW.__anext__ stores the
|
||||
# assembled response on logging_obj; the outer consumer
|
||||
# (ProxyLogging._fire_deferred_stream_logging) fires the
|
||||
# closure after the full streaming pipeline finishes.
|
||||
# The closure runs non-apply_guardrail hooks on the
|
||||
# assembled response, then fires success logging.
|
||||
# Only for CustomStreamWrapper — raw async generators from
|
||||
# passthrough routes bypass CSW and would orphan the closure.
|
||||
from litellm.litellm_core_utils.streaming_handler import (
|
||||
CustomStreamWrapper,
|
||||
)
|
||||
|
||||
if _post_call_guardrails_active and isinstance(response, CustomStreamWrapper):
|
||||
# Intentionally a live reference (not a copy) — mirrors
|
||||
# ProxyLogging.post_call_success_hook which also mutates
|
||||
# data["guardrail_to_apply"] during iteration.
|
||||
_captured_data: Final = self.data
|
||||
_captured_user_api_key_dict: Final = user_api_key_dict
|
||||
_captured_logging_obj: Final = logging_obj
|
||||
|
||||
async def _on_deferred_stream_complete(assembled_response: object, cache_hit: object) -> None:
|
||||
await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails(
|
||||
captured_data=_captured_data,
|
||||
captured_user_api_key_dict=_captured_user_api_key_dict,
|
||||
captured_logging_obj=_captured_logging_obj,
|
||||
assembled_response=assembled_response,
|
||||
cache_hit=cache_hit,
|
||||
)
|
||||
|
||||
logging_obj._on_deferred_stream_complete = _on_deferred_stream_complete
|
||||
elif (
|
||||
_post_call_guardrails_active
|
||||
and route_type in ("anthropic_messages", "aresponses")
|
||||
and self._is_streaming_response(response)
|
||||
):
|
||||
from litellm.litellm_core_utils.logging_worker import (
|
||||
GLOBAL_LOGGING_WORKER,
|
||||
if _post_call_guardrails_active:
|
||||
self._arm_deferred_stream_dispatch(
|
||||
response=response,
|
||||
route_type=route_type,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
async def _on_deferred_native_stream_complete(
|
||||
logging_coroutine: Coroutine[object, object, object],
|
||||
) -> None:
|
||||
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(async_coroutine=logging_coroutine)
|
||||
|
||||
logging_obj._on_deferred_stream_complete = _on_deferred_native_stream_complete
|
||||
|
||||
if route_type == "allm_passthrough_route":
|
||||
# Check if response is an async generator
|
||||
if self._is_streaming_response(response):
|
||||
|
|
@ -3057,6 +3018,87 @@ class ProxyBaseLLMRequestProcessing:
|
|||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error firing deferred logging: %s", e)
|
||||
|
||||
def _arm_deferred_stream_dispatch(
|
||||
self,
|
||||
response: object,
|
||||
route_type: str,
|
||||
user_api_key_dict: "UserAPIKeyAuth",
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> None:
|
||||
"""
|
||||
Streaming with post-call guardrails active: attach a closure that
|
||||
ProxyLogging._fire_deferred_stream_logging fires after all guardrail
|
||||
end-of-stream blocks complete, so the spend log sees
|
||||
guardrail_information.
|
||||
|
||||
Three closure shapes, matching who owns logging for the stream:
|
||||
- CustomStreamWrapper (chat completions) stores
|
||||
(assembled_response, cache_hit); the closure also runs
|
||||
non-apply_guardrail post-call hooks via
|
||||
_run_deferred_stream_guardrails.
|
||||
- Bridged /v1/responses (LiteLLMCompletionStreamingIterator) shares
|
||||
its inner CustomStreamWrapper's logging_obj, so it stores the same
|
||||
(assembled_response, cache_hit) shape; the closure only dispatches
|
||||
success logging, matching the route's pre-existing hook surface.
|
||||
- Native anthropic_messages/aresponses iterators store a single
|
||||
ready-made logging coroutine to enqueue.
|
||||
|
||||
Raw async generators from passthrough routes bypass all three and
|
||||
would orphan the closure, so they are not armed here.
|
||||
"""
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
|
||||
if isinstance(response, CustomStreamWrapper):
|
||||
# Intentionally a live reference (not a copy) — mirrors
|
||||
# ProxyLogging.post_call_success_hook which also mutates
|
||||
# data["guardrail_to_apply"] during iteration.
|
||||
_captured_data: Final = self.data
|
||||
_captured_user_api_key_dict: Final = user_api_key_dict
|
||||
_captured_logging_obj: Final = logging_obj
|
||||
|
||||
async def _on_deferred_stream_complete(assembled_response: object, cache_hit: object) -> None:
|
||||
await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails(
|
||||
captured_data=_captured_data,
|
||||
captured_user_api_key_dict=_captured_user_api_key_dict,
|
||||
captured_logging_obj=_captured_logging_obj,
|
||||
assembled_response=assembled_response,
|
||||
cache_hit=cache_hit,
|
||||
)
|
||||
|
||||
logging_obj._on_deferred_stream_complete = _on_deferred_stream_complete
|
||||
return
|
||||
|
||||
if route_type not in ("anthropic_messages", "aresponses") or not self._is_streaming_response(response):
|
||||
return
|
||||
|
||||
from litellm.responses.litellm_completion_transformation.streaming_iterator import (
|
||||
LiteLLMCompletionStreamingIterator,
|
||||
)
|
||||
|
||||
if isinstance(response, LiteLLMCompletionStreamingIterator):
|
||||
_captured_bridge_logging_obj: Final = logging_obj
|
||||
|
||||
async def _on_deferred_bridged_stream_complete(assembled_response: object, cache_hit: object) -> None:
|
||||
await _as_success_dispatcher(_captured_bridge_logging_obj).dispatch_success_handlers(
|
||||
assembled_response,
|
||||
cache_hit=cache_hit,
|
||||
start_time=None,
|
||||
end_time=None,
|
||||
prefer_async_handlers=True,
|
||||
)
|
||||
|
||||
logging_obj._on_deferred_stream_complete = _on_deferred_bridged_stream_complete
|
||||
return
|
||||
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
|
||||
async def _on_deferred_native_stream_complete(
|
||||
logging_coroutine: Coroutine[object, object, object],
|
||||
) -> None:
|
||||
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(async_coroutine=logging_coroutine)
|
||||
|
||||
logging_obj._on_deferred_stream_complete = _on_deferred_native_stream_complete
|
||||
|
||||
@staticmethod
|
||||
async def _run_deferred_stream_guardrails(
|
||||
captured_data: dict,
|
||||
|
|
|
|||
|
|
@ -65,6 +65,19 @@ class PassThroughStreamingHandler:
|
|||
route_streaming_logging or PassThroughStreamingHandler._route_streaming_logging_to_handler
|
||||
)
|
||||
raw_bytes: Final[list[bytes]] = []
|
||||
|
||||
def _build_logging_coroutine() -> Coroutine[None, None, None]:
|
||||
return resolved_route_streaming_logging(
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
passthrough_success_handler_obj=passthrough_success_handler_obj,
|
||||
url_route=url_route,
|
||||
request_body=request_body or {},
|
||||
endpoint_type=endpoint_type,
|
||||
start_time=start_time,
|
||||
raw_bytes=raw_bytes,
|
||||
end_time=datetime.now(),
|
||||
)
|
||||
|
||||
logging_scheduled = False
|
||||
model_name: Final = PassThroughStreamingHandler._extract_model_for_cost_injection(
|
||||
request_body=request_body,
|
||||
|
|
@ -114,6 +127,21 @@ class PassThroughStreamingHandler:
|
|||
)
|
||||
if pending:
|
||||
yield pending
|
||||
# Stream completed cleanly. When the proxy armed deferred
|
||||
# dispatch (post-call guardrails active), park the logging
|
||||
# coroutine on logging_obj instead of enqueueing now, so
|
||||
# ProxyLogging._fire_deferred_stream_logging fires it after
|
||||
# guardrail end-of-stream blocks populate guardrail_information.
|
||||
# Disconnect/exception paths skip this and fall through to the
|
||||
# immediate enqueue in ``finally`` to keep partial billing
|
||||
# (LIT-2642).
|
||||
if (
|
||||
getattr(litellm_logging_obj, "_on_deferred_stream_complete", None) is not None
|
||||
and raw_bytes
|
||||
and response.status_code < 400
|
||||
):
|
||||
logging_scheduled = True
|
||||
litellm_logging_obj._deferred_stream_complete_args = (_build_logging_coroutine(),)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("Error in chunk_processor: %s", e)
|
||||
raise
|
||||
|
|
@ -128,18 +156,7 @@ class PassThroughStreamingHandler:
|
|||
if not logging_scheduled and raw_bytes and response.status_code < 400:
|
||||
logging_scheduled = True
|
||||
try:
|
||||
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
|
||||
async_coroutine=resolved_route_streaming_logging(
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
passthrough_success_handler_obj=passthrough_success_handler_obj,
|
||||
url_route=url_route,
|
||||
request_body=request_body or {},
|
||||
endpoint_type=endpoint_type,
|
||||
start_time=start_time,
|
||||
raw_bytes=raw_bytes,
|
||||
end_time=datetime.now(),
|
||||
)
|
||||
)
|
||||
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(async_coroutine=_build_logging_coroutine())
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("Error scheduling chunk_processor logging: %s", e)
|
||||
|
||||
|
|
|
|||
|
|
@ -1297,3 +1297,126 @@ class TestResponsesIteratorDeferredLogging:
|
|||
assert len(created) == 1
|
||||
await created[0]
|
||||
assert recorded["dispatched"] is True
|
||||
|
||||
|
||||
class TestArmDeferredStreamDispatch:
|
||||
"""Regression for PR #38722: the closure shape armed on logging_obj must
|
||||
match the args the stream's logging owner stores. Bridged /v1/responses
|
||||
(LiteLLMCompletionStreamingIterator) shares its inner CustomStreamWrapper's
|
||||
logging_obj, which stores (assembled_response, cache_hit); arming the
|
||||
single-coroutine native closure there made _fire_deferred_stream_logging
|
||||
raise TypeError inside the streaming hook, leaking an in-stream 500 error
|
||||
frame on every streamed /v1/responses request."""
|
||||
|
||||
def _processor(self):
|
||||
return ProxyBaseLLMRequestProcessing(data={"model": "gpt-test"})
|
||||
|
||||
def _dispatch_recording_logging_obj(self):
|
||||
recorded = {}
|
||||
|
||||
async def dispatch_success_handlers(
|
||||
result=None, start_time=None, end_time=None, cache_hit=None, prefer_async_handlers=False
|
||||
):
|
||||
recorded["result"] = result
|
||||
recorded["cache_hit"] = cache_hit
|
||||
recorded["prefer_async_handlers"] = prefer_async_handlers
|
||||
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.dispatch_success_handlers = dispatch_success_handlers
|
||||
logging_obj._on_deferred_stream_complete = None
|
||||
logging_obj._deferred_stream_complete_args = None
|
||||
return logging_obj, recorded
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bridged_responses_iterator_gets_csw_arg_shape(self):
|
||||
from litellm.responses.litellm_completion_transformation.streaming_iterator import (
|
||||
LiteLLMCompletionStreamingIterator,
|
||||
)
|
||||
|
||||
logging_obj, recorded = self._dispatch_recording_logging_obj()
|
||||
bridged = object.__new__(LiteLLMCompletionStreamingIterator)
|
||||
|
||||
self._processor()._arm_deferred_stream_dispatch(
|
||||
response=bridged,
|
||||
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
|
||||
|
||||
logging_obj, _ = self._dispatch_recording_logging_obj()
|
||||
|
||||
async def _agen():
|
||||
yield b"x"
|
||||
|
||||
self._processor()._arm_deferred_stream_dispatch(
|
||||
response=_agen(),
|
||||
route_type="anthropic_messages",
|
||||
user_api_key_dict=MagicMock(),
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
closure = logging_obj._on_deferred_stream_complete
|
||||
assert closure is not None
|
||||
|
||||
async def _logging_coroutine():
|
||||
return None
|
||||
|
||||
coro = _logging_coroutine()
|
||||
with patch.object( # test-quality-ok: GLOBAL_LOGGING_WORKER is a process-global singleton with no injection seam
|
||||
GLOBAL_LOGGING_WORKER, "ensure_initialized_and_enqueue"
|
||||
) as mock_enqueue:
|
||||
await closure(coro)
|
||||
mock_enqueue.assert_called_once_with(async_coroutine=coro)
|
||||
coro.close()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_csw_closure_routes_through_deferred_stream_guardrails(self, monkeypatch):
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
|
||||
logging_obj, recorded = self._dispatch_recording_logging_obj()
|
||||
csw = object.__new__(CustomStreamWrapper)
|
||||
processor = self._processor()
|
||||
|
||||
monkeypatch.setattr( # test-quality-ok: empty the process-global callback registry so no ambient guardrail runs
|
||||
litellm, "callbacks", []
|
||||
)
|
||||
processor._arm_deferred_stream_dispatch(
|
||||
response=csw,
|
||||
route_type="acompletion",
|
||||
user_api_key_dict=MagicMock(),
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
assembled = object()
|
||||
await logging_obj._on_deferred_stream_complete(assembled, False)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert recorded["result"] is assembled
|
||||
assert recorded["cache_hit"] is False
|
||||
assert recorded["prefer_async_handlers"] is True
|
||||
|
||||
def test_non_native_route_generator_not_armed(self):
|
||||
logging_obj, _ = self._dispatch_recording_logging_obj()
|
||||
|
||||
async def _agen():
|
||||
yield b"x"
|
||||
|
||||
self._processor()._arm_deferred_stream_dispatch(
|
||||
response=_agen(),
|
||||
route_type="acompletion",
|
||||
user_api_key_dict=MagicMock(),
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
assert logging_obj._on_deferred_stream_complete is None
|
||||
|
|
|
|||
|
|
@ -28,12 +28,21 @@ def _make_streaming_response(chunks):
|
|||
return mock
|
||||
|
||||
|
||||
def _unarmed_logging_obj():
|
||||
"""Real Logging objects only carry _on_deferred_stream_complete when the
|
||||
proxy arms deferred dispatch; a bare MagicMock's auto-attribute is truthy
|
||||
and would spuriously trigger the deferral branch."""
|
||||
obj = MagicMock()
|
||||
obj._on_deferred_stream_complete = None
|
||||
return obj
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chunk_processor_logs_on_normal_completion():
|
||||
chunks = [b"chunk-1", b"chunk-2", b"chunk-3"]
|
||||
response = _make_streaming_response(chunks)
|
||||
|
||||
mock_logging_obj = MagicMock()
|
||||
mock_logging_obj = _unarmed_logging_obj()
|
||||
mock_passthrough_handler = MagicMock()
|
||||
|
||||
with patch.object(
|
||||
|
|
@ -66,7 +75,7 @@ async def test_chunk_processor_logs_on_client_disconnect():
|
|||
chunks = [b"event-1", b"event-2", b"event-3"]
|
||||
response = _make_streaming_response(chunks)
|
||||
|
||||
mock_logging_obj = MagicMock()
|
||||
mock_logging_obj = _unarmed_logging_obj()
|
||||
mock_passthrough_handler = MagicMock()
|
||||
|
||||
with patch.object(
|
||||
|
|
@ -104,7 +113,7 @@ async def test_chunk_processor_does_not_schedule_success_logging_for_upstream_er
|
|||
response = _make_streaming_response(chunks)
|
||||
response.status_code = 403
|
||||
|
||||
mock_logging_obj = MagicMock()
|
||||
mock_logging_obj = _unarmed_logging_obj()
|
||||
mock_passthrough_handler = MagicMock()
|
||||
|
||||
with patch.object(
|
||||
|
|
@ -134,7 +143,7 @@ async def test_chunk_processor_does_not_schedule_success_logging_for_upstream_er
|
|||
async def test_chunk_processor_does_not_schedule_logging_when_no_chunks():
|
||||
response = _make_streaming_response([])
|
||||
|
||||
mock_logging_obj = MagicMock()
|
||||
mock_logging_obj = _unarmed_logging_obj()
|
||||
mock_passthrough_handler = MagicMock()
|
||||
|
||||
with patch.object(
|
||||
|
|
@ -189,7 +198,7 @@ async def test_chunk_processor_routes_logging_through_logging_worker():
|
|||
async for chunk in PassThroughStreamingHandler.chunk_processor(
|
||||
response=response,
|
||||
request_body={"model": "claude-3-haiku"},
|
||||
litellm_logging_obj=MagicMock(),
|
||||
litellm_logging_obj=_unarmed_logging_obj(),
|
||||
endpoint_type=EndpointType.GENERIC,
|
||||
start_time=datetime.now(),
|
||||
passthrough_success_handler_obj=MagicMock(),
|
||||
|
|
@ -230,7 +239,7 @@ async def test_chunk_processor_routes_logging_through_logging_worker_on_disconne
|
|||
gen = PassThroughStreamingHandler.chunk_processor(
|
||||
response=response,
|
||||
request_body={"model": "claude-3-haiku"},
|
||||
litellm_logging_obj=MagicMock(),
|
||||
litellm_logging_obj=_unarmed_logging_obj(),
|
||||
endpoint_type=EndpointType.GENERIC,
|
||||
start_time=datetime.now(),
|
||||
passthrough_success_handler_obj=MagicMock(),
|
||||
|
|
@ -246,7 +255,7 @@ async def test_chunk_processor_routes_logging_through_logging_worker_on_disconne
|
|||
def _logging_obj_with_write_once_cst():
|
||||
"""Build a MagicMock that mirrors the real Logging behavior: _update_completion_start_time
|
||||
latches self.completion_start_time so the write-once guard actually latches."""
|
||||
obj = MagicMock()
|
||||
obj = _unarmed_logging_obj()
|
||||
obj.completion_start_time = None
|
||||
|
||||
def _update(*, completion_start_time):
|
||||
|
|
@ -301,7 +310,7 @@ async def test_chunk_processor_does_not_reset_completion_start_time_on_later_chu
|
|||
response = _make_streaming_response(chunks)
|
||||
|
||||
real_first = datetime(2020, 1, 1, 0, 0, 0)
|
||||
mock_logging_obj = MagicMock()
|
||||
mock_logging_obj = _unarmed_logging_obj()
|
||||
# Simulate first-chunk stamp having already landed (e.g. under contention or a
|
||||
# prior wrapper that already set it): later chunks must be no-ops.
|
||||
mock_logging_obj.completion_start_time = real_first
|
||||
|
|
@ -387,7 +396,7 @@ async def _collect_openai_passthrough_chunks(chunks, endpoint_type):
|
|||
async for chunk in PassThroughStreamingHandler.chunk_processor(
|
||||
response=response,
|
||||
request_body={"model": "gpt-4o-mini", "stream": True},
|
||||
litellm_logging_obj=MagicMock(),
|
||||
litellm_logging_obj=_unarmed_logging_obj(),
|
||||
endpoint_type=endpoint_type,
|
||||
start_time=datetime.now(),
|
||||
passthrough_success_handler_obj=MagicMock(),
|
||||
|
|
@ -517,3 +526,109 @@ def test_convert_raw_bytes_survives_truncated_multibyte_sequence():
|
|||
lines = PassThroughStreamingHandler._convert_raw_bytes_to_str_lines(raw_bytes)
|
||||
|
||||
assert any('"type": "message_delta"' in line for line in lines)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chunk_processor_defers_logging_until_fire_when_armed():
|
||||
"""Regression for PR #38722: native /v1/messages streams route through
|
||||
chunk_processor, which enqueued the spend log the moment the stream ended,
|
||||
racing the guardrail end-of-stream scan and logging
|
||||
guardrail_information as null. With deferred dispatch armed, the completed
|
||||
stream must park the logging coroutine on logging_obj and only enqueue it
|
||||
when ProxyLogging._fire_deferred_stream_logging fires after the scan."""
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
||||
chunks = [b"event-1", b"event-2"]
|
||||
response = _make_streaming_response(chunks)
|
||||
|
||||
logging_obj = _unarmed_logging_obj()
|
||||
logging_obj._deferred_stream_complete_args = None
|
||||
|
||||
enqueued = []
|
||||
|
||||
def _capture(async_coroutine):
|
||||
enqueued.append(async_coroutine)
|
||||
async_coroutine.close()
|
||||
|
||||
with patch.object( # test-quality-ok: GLOBAL_LOGGING_WORKER is a process-global singleton with no injection seam
|
||||
GLOBAL_LOGGING_WORKER,
|
||||
"ensure_initialized_and_enqueue",
|
||||
side_effect=_capture,
|
||||
) as mock_enqueue:
|
||||
gen = PassThroughStreamingHandler.chunk_processor(
|
||||
response=response,
|
||||
request_body={"model": "claude-3-haiku"},
|
||||
litellm_logging_obj=logging_obj,
|
||||
endpoint_type=EndpointType.ANTHROPIC,
|
||||
start_time=datetime.now(),
|
||||
passthrough_success_handler_obj=MagicMock(),
|
||||
url_route="/v1/messages",
|
||||
route_streaming_logging=AsyncMock(),
|
||||
)
|
||||
ProxyBaseLLMRequestProcessing(data={})._arm_deferred_stream_dispatch(
|
||||
response=gen,
|
||||
route_type="anthropic_messages",
|
||||
user_api_key_dict=MagicMock(),
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
received = []
|
||||
async for chunk in gen:
|
||||
received.append(chunk)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert received == chunks
|
||||
mock_enqueue.assert_not_called()
|
||||
parked = logging_obj._deferred_stream_complete_args
|
||||
assert isinstance(parked, tuple) and len(parked) == 1
|
||||
assert asyncio.iscoroutine(parked[0])
|
||||
|
||||
ProxyLogging._fire_deferred_stream_logging({"litellm_logging_obj": logging_obj})
|
||||
await asyncio.sleep(0)
|
||||
|
||||
mock_enqueue.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chunk_processor_enqueues_immediately_on_disconnect_even_when_armed():
|
||||
"""Client disconnects never reach _fire_deferred_stream_logging, so parking
|
||||
the coroutine there would lose the partial-usage spend log (LIT-2642); the
|
||||
disconnect path must keep enqueueing immediately."""
|
||||
chunks = [b"event-1", b"event-2", b"event-3"]
|
||||
response = _make_streaming_response(chunks)
|
||||
|
||||
logging_obj = _unarmed_logging_obj()
|
||||
|
||||
async def _armed_closure(logging_coroutine):
|
||||
raise AssertionError("deferred closure must not fire on disconnect")
|
||||
|
||||
logging_obj._on_deferred_stream_complete = _armed_closure
|
||||
logging_obj._deferred_stream_complete_args = None
|
||||
|
||||
enqueued = []
|
||||
|
||||
def _capture(async_coroutine):
|
||||
enqueued.append(async_coroutine)
|
||||
async_coroutine.close()
|
||||
|
||||
with patch.object( # test-quality-ok: GLOBAL_LOGGING_WORKER is a process-global singleton with no injection seam
|
||||
GLOBAL_LOGGING_WORKER,
|
||||
"ensure_initialized_and_enqueue",
|
||||
side_effect=_capture,
|
||||
) as mock_enqueue:
|
||||
gen = PassThroughStreamingHandler.chunk_processor(
|
||||
response=response,
|
||||
request_body={"model": "claude-3-haiku"},
|
||||
litellm_logging_obj=logging_obj,
|
||||
endpoint_type=EndpointType.ANTHROPIC,
|
||||
start_time=datetime.now(),
|
||||
passthrough_success_handler_obj=MagicMock(),
|
||||
url_route="/v1/messages",
|
||||
route_streaming_logging=AsyncMock(),
|
||||
)
|
||||
await gen.__anext__()
|
||||
await gen.aclose()
|
||||
|
||||
mock_enqueue.assert_called_once()
|
||||
assert logging_obj._deferred_stream_complete_args is None
|
||||
|
|
|
|||
|
|
@ -1656,7 +1656,7 @@ class TestCommonRequestProcessingHelpers:
|
|||
assert payload["error"]["message"] == "MCP request blocked: no rewritable argument field present"
|
||||
assert payload["error"]["provider_specific_fields"]["error"]["code"] == "panw_prisma_airs_blocked"
|
||||
|
||||
async def testserialize_http_exception_detail_helper(self):
|
||||
async def test_serialize_http_exception_detail_helper(self):
|
||||
"""Direct unit coverage for the L1 helper across all branches."""
|
||||
from litellm.proxy.common_request_processing import (
|
||||
serialize_http_exception_detail,
|
||||
|
|
|
|||
|
|
@ -346,14 +346,14 @@ async def test_post_call_stream_guardrail_keeps_own_iterator_on_chat_completions
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unified_guardrail_iterator_accepts_explicit_guardrail(monkeypatch):
|
||||
async def test_unified_guardrail_iterator_accepts_explicit_guardrail():
|
||||
"""
|
||||
The dispatch passes each guardrail explicitly instead of through a shared
|
||||
request_data key, so chaining two unified-routed guardrails cannot drop
|
||||
all but the last one.
|
||||
all but the last one. The block fires after the deltas were already
|
||||
flushed to the client, so it surfaces as a trailing in-stream error frame
|
||||
rather than a raised HTTPException.
|
||||
"""
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy.utils import unified_guardrail
|
||||
|
||||
guardrail = _content_filter_guardrail("BLOCK")
|
||||
|
|
@ -367,14 +367,19 @@ async def test_unified_guardrail_iterator_accepts_explicit_guardrail(monkeypatch
|
|||
for chunk in _anthropic_stream_chunks(["the", " zebra runs"]):
|
||||
yield chunk
|
||||
|
||||
with pytest.raises(HTTPException):
|
||||
async for _ in unified_guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-1234", request_route="/v1/messages"),
|
||||
response=fake_stream(),
|
||||
request_data=request_data,
|
||||
guardrail_to_apply=guardrail,
|
||||
):
|
||||
pass
|
||||
delivered = []
|
||||
async for item in unified_guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-1234", request_route="/v1/messages"),
|
||||
response=fake_stream(),
|
||||
request_data=request_data,
|
||||
guardrail_to_apply=guardrail,
|
||||
):
|
||||
delivered.append(item)
|
||||
|
||||
raw = b"".join(c for c in delivered if isinstance(c, bytes)).decode()
|
||||
assert "event: error" in raw
|
||||
assert "guardrail_error" in raw
|
||||
assert raw.index("guardrail_error") > raw.index(" zebra runs")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue