mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
fix(guardrails): surface a during_call block without waiting for slower guardrails
This commit is contained in:
parent
9f177580d7
commit
ebce237785
2 changed files with 106 additions and 13 deletions
|
|
@ -415,6 +415,37 @@ def _exception_changes_request_flow(exc: BaseException) -> bool:
|
|||
return isinstance(exc, (SensitiveDataRouteException, ModifyResponseException))
|
||||
|
||||
|
||||
async def _capture_flow_change(coro: Awaitable[None]) -> BaseException | None:
|
||||
"""
|
||||
Await `coro`, returning a reroute/passthrough exception as a value instead of
|
||||
raising it and letting a blocking exception propagate.
|
||||
|
||||
This is what lets `asyncio.wait(..., FIRST_EXCEPTION)` return on the first
|
||||
*block* while still awaiting every guardrail when the only exceptions so far
|
||||
change the request flow: a slower block must still win over a faster reroute.
|
||||
"""
|
||||
try:
|
||||
await coro
|
||||
except Exception as exc:
|
||||
if _exception_changes_request_flow(exc):
|
||||
return exc
|
||||
raise
|
||||
return None
|
||||
|
||||
|
||||
async def _cancel_and_drain_guardrail_tasks(tasks: tuple["asyncio.Task[BaseException | None]", ...]) -> None:
|
||||
"""
|
||||
Cancel whatever is still running and await every task, so a guardrail
|
||||
abandoned once another one blocked is never left as an unobserved background
|
||||
task. Cancellation is delivered at the abandoned guardrail's next await, so
|
||||
this does not wait for it to finish its scan.
|
||||
"""
|
||||
for task in tasks:
|
||||
if not task.done():
|
||||
task.cancel()
|
||||
await asyncio.gather(*tasks, return_exceptions=True)
|
||||
|
||||
|
||||
def _policy_state_metadata(data: Mapping[str, object]) -> Mapping[str, object]:
|
||||
"""
|
||||
Return the metadata bucket the policy engine wrote its pipeline state into.
|
||||
|
|
@ -2058,6 +2089,13 @@ class ProxyLogging:
|
|||
except SensitiveDataRouteException:
|
||||
status = "intervened"
|
||||
raise
|
||||
except asyncio.CancelledError:
|
||||
# CancelledError is a BaseException, so without this branch the
|
||||
# `finally` below would record a guardrail abandoned mid-flight
|
||||
# (a sibling blocked, or the racing provider call raised first) as
|
||||
# a phantom `success` sample.
|
||||
status = "cancelled" # rebind-ok: the status this `finally` reports is settled per-branch
|
||||
raise
|
||||
except Exception as e:
|
||||
status = "error"
|
||||
error_type = type(e).__name__
|
||||
|
|
@ -2253,11 +2291,17 @@ class ProxyLogging:
|
|||
Each per-guardrail coroutine sets ``guardrail_to_apply`` immediately
|
||||
before awaiting, and the unified hook pops it before its first
|
||||
suspension point, so concurrent guardrails never race on that key.
|
||||
Every guardrail is awaited to completion (``return_exceptions=True``)
|
||||
so a raise by one never leaves the others running as unobserved
|
||||
background tasks. A guardrail that blocks (any exception other than a
|
||||
reroute or passthrough) takes precedence over one that only changes the
|
||||
request flow, so a fast reroute can never let a slower block be bypassed.
|
||||
|
||||
A guardrail that blocks (any exception other than a reroute or
|
||||
passthrough) is surfaced as soon as it raises, without waiting for
|
||||
slower siblings: ``ProxyBaseLLMRequestProcessing`` races this hook
|
||||
against the provider call and cancels the loser, so every millisecond
|
||||
spent waiting once a block is known is a millisecond the provider call
|
||||
has to finish and bill for a request that is being rejected anyway. A
|
||||
reroute does *not* return early -- a slower block still takes precedence
|
||||
over a faster reroute, so a crafted input can never dodge a block.
|
||||
Whatever is still running once a block wins is cancelled and awaited, so
|
||||
no guardrail is left as an unobserved background task.
|
||||
"""
|
||||
# Fast path: skip the entire guardrail scan when no CustomGuardrail
|
||||
# callbacks are registered. Saves per-request iteration over
|
||||
|
|
@ -2326,16 +2370,28 @@ class ProxyLogging:
|
|||
# Add task to list for parallel execution
|
||||
guardrail_tasks.append(_run_one(callback))
|
||||
|
||||
# Step 2: Run all guardrail tasks in parallel
|
||||
# Step 2: Run all guardrail tasks in parallel, stopping at the first block
|
||||
if not guardrail_tasks:
|
||||
return data
|
||||
results: Final = await asyncio.gather(*guardrail_tasks, return_exceptions=True)
|
||||
raised: Final = tuple(result for result in results if isinstance(result, BaseException))
|
||||
blocking: Final = next((exc for exc in raised if not _exception_changes_request_flow(exc)), None)
|
||||
if blocking is not None:
|
||||
raise blocking
|
||||
if raised:
|
||||
raise raised[0]
|
||||
tasks: Final = tuple(asyncio.create_task(_capture_flow_change(coro)) for coro in guardrail_tasks)
|
||||
try:
|
||||
done, _pending = await asyncio.wait(tasks, return_when=asyncio.FIRST_EXCEPTION)
|
||||
# Iterate `tasks`, not the `done` set, so registration order decides
|
||||
# which exception wins when several land in the same pass.
|
||||
blocking: Final = tuple(
|
||||
exc for exc in (task.exception() for task in tasks if task in done) if exc is not None
|
||||
)
|
||||
reroutes: Final = tuple(
|
||||
exc
|
||||
for exc in (task.result() for task in tasks if task in done and task.exception() is None)
|
||||
if exc is not None
|
||||
)
|
||||
finally:
|
||||
await _cancel_and_drain_guardrail_tasks(tasks)
|
||||
if blocking:
|
||||
raise blocking[0]
|
||||
if reroutes:
|
||||
raise reroutes[0]
|
||||
|
||||
return data
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import asyncio
|
||||
import datetime as real_datetime
|
||||
import smtplib
|
||||
import time
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -2061,3 +2062,39 @@ async def test_during_call_hook_block_wins_over_reroute(install_guardrails):
|
|||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "blocked by guardrail" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_during_call_hook_block_does_not_wait_for_slower_guardrail(install_guardrails):
|
||||
"""A block is surfaced immediately: this hook races the provider call, so holding the
|
||||
rejection for a slow sibling gives the provider call time to complete and bill."""
|
||||
execution_order = []
|
||||
install_guardrails(
|
||||
[
|
||||
_UnifiedDuringCallGuardrail(
|
||||
"blocker",
|
||||
execution_order,
|
||||
sleep=0,
|
||||
raise_exc=HTTPException(status_code=400, detail="blocked by guardrail"),
|
||||
),
|
||||
_UnifiedDuringCallGuardrail("slowpoke", execution_order, sleep=5),
|
||||
]
|
||||
)
|
||||
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
|
||||
|
||||
start = time.monotonic()
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await proxy_logging_obj.during_call_hook(
|
||||
data=_during_call_request_data(),
|
||||
user_api_key_dict=_during_call_key(),
|
||||
call_type="acompletion",
|
||||
)
|
||||
elapsed = time.monotonic() - start
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert elapsed < 1, f"block waited {elapsed:.2f}s for the slower guardrail"
|
||||
# The abandoned guardrail is cancelled, not left running as an unobserved
|
||||
# background task that finishes its scan after the request was rejected.
|
||||
assert "slowpoke_start" in execution_order
|
||||
await asyncio.sleep(0.05)
|
||||
assert "slowpoke_end" not in execution_order
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue