perf(guardrails): build fire_and_forget payloads only when a slot is free

BackgroundDispatcher.dispatch now takes a prepare callable and calls it only after the in-flight cap check passes. A call dropped by fire_and_forget_max_inflight no longer serializes its payload on the request path, which matters most for large or image-heavy requests while the endpoint is slow

A payload that cannot be serialized is still recorded as not_run before anything is dispatched
This commit is contained in:
Caduri Katzav 2026-10-03 15:48:31 +03:00
parent f7e934677b
commit 09922b74cd
4 changed files with 42 additions and 14 deletions

View file

@ -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

View file

@ -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(

View file

@ -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()

View file

@ -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)