mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
f7e934677b
commit
09922b74cd
4 changed files with 42 additions and 14 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue