diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/background_dispatch.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/background_dispatch.py index 247ff2a9b47..ea48c1c22f2 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/background_dispatch.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/background_dispatch.py @@ -74,7 +74,7 @@ class BackgroundDispatcher: def dropped_count(self) -> int: return self._dropped - def dispatch(self, run: Callable[[], Awaitable[None]], *, context: str) -> bool: + def dispatch(self, prepare: Callable[[], Callable[[], Awaitable[object]]], *, context: str) -> bool: if len(self._pending) >= self._max_inflight: self._dropped += 1 if self._dropped % _DROP_LOG_INTERVAL == 1: @@ -89,12 +89,13 @@ class BackgroundDispatcher: ) return False + run: Final = prepare() task: Final = contextvars.Context().run(asyncio.create_task, self._run_logging_failures(run, context=context)) self._pending.add(task) task.add_done_callback(self._pending.discard) return True - async def _run_logging_failures(self, run: Callable[[], Awaitable[None]], *, context: str) -> None: + async def _run_logging_failures(self, run: Callable[[], Awaitable[object]], *, context: str) -> None: try: await run() except Exception as e: # noqa: BLE001 # a detached task has no caller to raise into diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py index ee1b8a72973..b50e9adbb0c 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py @@ -7,7 +7,7 @@ import fnmatch import os -from collections.abc import Mapping, Sequence +from collections.abc import Awaitable, Callable, Mapping, Sequence from typing import TYPE_CHECKING, Any, Final, Literal, Optional import httpx @@ -442,14 +442,18 @@ class GenericGuardrailAPI(CustomGuardrail): input_type: Literal["request", "response"], logging_obj: Optional["LiteLLMLoggingObj"], ) -> bool: - payload: Final = guardrail_request.model_dump(mode="json") - headers: Final = self._build_request_headers() timeout: Final = FIRE_AND_FORGET_POST_TIMEOUT_SECONDS if self.timeout is None else self.timeout - async def _post() -> None: - await self.async_handler.post(url=self.api_base, json=payload, headers=headers, timeout=timeout) + def _prepare() -> Callable[[], Awaitable[None]]: + payload: Final = guardrail_request.model_dump(mode="json") + headers: Final = self._build_request_headers() - return self._dispatcher.dispatch(_post, context=_call_context(input_type, logging_obj)) + async def _post() -> None: + await self.async_handler.post(url=self.api_base, json=payload, headers=headers, timeout=timeout) + + return _post + + return self._dispatcher.dispatch(_prepare, context=_call_context(input_type, logging_obj)) @log_guardrail_information async def apply_guardrail( diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_background_dispatch.py b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_background_dispatch.py index 30bb7e353c7..9f660dae6df 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_background_dispatch.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_background_dispatch.py @@ -1,6 +1,7 @@ import asyncio import contextvars -from collections.abc import Callable +from collections.abc import Awaitable, Callable +from functools import partial from typing import Final import pytest @@ -49,15 +50,21 @@ def test_a_dispatcher_needs_room_for_at_least_one_call() -> None: BackgroundDispatcher(guardrail_name="g", max_inflight=0) -async def test_calls_beyond_the_cap_are_dropped_counted_and_warned_once( +async def test_calls_beyond_the_cap_are_dropped_unprepared_counted_and_warned_once( warning_messages: Callable[[str], list[str]], ) -> None: gate: Final = asyncio.Event() dispatcher: Final = BackgroundDispatcher(guardrail_name="g", max_inflight=2) + prepared: Final[list[int]] = [] # mutable-ok: records which calls were prepared - dispatched: Final = [dispatcher.dispatch(gate.wait, context=f"call {i}") for i in range(5)] + def prepare(call: int) -> Callable[[], Awaitable[bool]]: + prepared.append(call) + return gate.wait + + dispatched: Final = [dispatcher.dispatch(partial(prepare, i), context=f"call {i}") for i in range(5)] assert (dispatched, dispatcher.pending_count, dispatcher.dropped_count) == ([True, True, False, False, False], 2, 3) + assert prepared == [0, 1] assert len(warning_messages("dropped")) == 1 gate.set() await dispatcher.wait_for_pending() @@ -70,7 +77,7 @@ async def test_a_finished_call_frees_its_slot() -> None: return None for _ in range(3): - assert dispatcher.dispatch(finish, context="call") is True + assert dispatcher.dispatch(lambda: finish, context="call") is True await dispatcher.wait_for_pending() assert (dispatcher.pending_count, dispatcher.dropped_count) == (0, 0) @@ -84,7 +91,7 @@ async def test_a_failing_call_is_logged_with_its_context_and_not_raised( async def fail() -> None: raise ConnectionError("refused") - dispatcher.dispatch(fail, context="input_type=response litellm_call_id=call-1") + dispatcher.dispatch(lambda: fail, context="input_type=response litellm_call_id=call-1") await dispatcher.wait_for_pending() assert warning_messages("call failed") == [ @@ -102,7 +109,7 @@ async def test_a_dispatched_call_does_not_see_the_request_context() -> None: token: Final = _request_scoped.set("request-1") try: - dispatcher.dispatch(record, context="call") + dispatcher.dispatch(lambda: record, context="call") finally: _request_scoped.reset(token) await dispatcher.wait_for_pending() diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_generic_guardrail_api.py b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_generic_guardrail_api.py index fce2c0845a6..909318966f1 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_generic_guardrail_api.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_generic_guardrail_api.py @@ -268,6 +268,22 @@ async def test_a_payload_that_cannot_be_serialized_is_recorded_as_not_dispatched assert _recorded_outcomes(request) == [("not_run", FIRE_AND_FORGET_NOT_DISPATCHED_REASON)] +async def test_a_call_dropped_by_the_cap_never_builds_its_payload() -> None: + endpoint: Final = _Endpoint() + guardrail, dispatcher = _fire_and_forget( + endpoint, max_inflight=1, additional_provider_specific_params={"bad": object()} + ) + gate: Final = asyncio.Event() + dispatcher.dispatch(lambda: gate.wait, context="occupies the only slot") + request: Final = _request_data() + + await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data=request, input_type="request") + gate.set() + await dispatcher.wait_for_pending() + + assert _recorded_outcomes(request) == [("not_run", FIRE_AND_FORGET_DROPPED_REASON)] + + async def test_dispatched_calls_are_recorded_as_success_and_dropped_ones_as_not_run() -> None: endpoint: Final = _Endpoint(gate_open=False) guardrail, dispatcher = _fire_and_forget(endpoint, max_inflight=1)