added guardrail retry logic

This commit is contained in:
shivam 2026-02-05 14:25:53 -08:00
parent 0649720f79
commit f2370bd3ba
8 changed files with 517 additions and 51 deletions

View file

@ -11,6 +11,10 @@ from typing import Any, AsyncGenerator, List, Optional, Union
from litellm._logging import verbose_proxy_logger
from litellm.caching.caching import DualCache
from litellm.proxy.guardrails.guardrail_retries import (
get_guardrail_retry_config,
run_guardrail_with_retries,
)
from litellm.cost_calculator import _infer_call_type
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_logger import CustomLogger
@ -91,10 +95,16 @@ class UnifiedLLMGuardrails(CustomLogger):
CallTypes(call_type)
]()
data = await endpoint_translation.process_input_messages(
data=data,
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=data.get("litellm_logging_obj"),
num_retries, retry_after = get_guardrail_retry_config(guardrail_to_apply)
data = await run_guardrail_with_retries(
coro_factory=lambda: endpoint_translation.process_input_messages(
data=data,
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=data.get("litellm_logging_obj"),
),
num_retries=num_retries,
retry_after=retry_after,
guardrail_name=guardrail_to_apply.guardrail_name,
)
# Add guardrail to applied guardrails header
@ -148,10 +158,16 @@ class UnifiedLLMGuardrails(CustomLogger):
CallTypes(call_type)
]()
return await endpoint_translation.process_input_messages(
data=data,
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=data.get("litellm_logging_obj"),
num_retries, retry_after = get_guardrail_retry_config(guardrail_to_apply)
return await run_guardrail_with_retries(
coro_factory=lambda: endpoint_translation.process_input_messages(
data=data,
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=data.get("litellm_logging_obj"),
),
num_retries=num_retries,
retry_after=retry_after,
guardrail_name=guardrail_to_apply.guardrail_name,
)
async def async_post_call_success_hook(
@ -213,11 +229,17 @@ class UnifiedLLMGuardrails(CustomLogger):
CallTypes(call_type)
]()
response = await endpoint_translation.process_output_response(
response=response, # type: ignore
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=data.get("litellm_logging_obj"),
user_api_key_dict=user_api_key_dict,
num_retries, retry_after = get_guardrail_retry_config(guardrail_to_apply)
response = await run_guardrail_with_retries(
coro_factory=lambda: endpoint_translation.process_output_response(
response=response, # type: ignore
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=data.get("litellm_logging_obj"),
user_api_key_dict=user_api_key_dict,
),
num_retries=num_retries,
retry_after=retry_after,
guardrail_name=guardrail_to_apply.guardrail_name,
)
# Add guardrail to applied guardrails header
add_guardrail_to_applied_guardrails_header(
@ -362,11 +384,19 @@ class UnifiedLLMGuardrails(CustomLogger):
CallTypes(call_type)
]()
await endpoint_translation.process_output_streaming_response(
responses_so_far=responses_so_far,
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=request_data.get("litellm_logging_obj"),
user_api_key_dict=user_api_key_dict,
num_retries, retry_after = get_guardrail_retry_config(
guardrail_to_apply
)
await run_guardrail_with_retries(
coro_factory=lambda: endpoint_translation.process_output_streaming_response(
responses_so_far=responses_so_far,
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=request_data.get("litellm_logging_obj"),
user_api_key_dict=user_api_key_dict,
),
num_retries=num_retries,
retry_after=retry_after,
guardrail_name=guardrail_to_apply.guardrail_name,
)
yield original_item
@ -388,9 +418,15 @@ class UnifiedLLMGuardrails(CustomLogger):
CallTypes(call_type)
]()
await endpoint_translation.process_output_streaming_response(
responses_so_far=responses_so_far,
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=request_data.get("litellm_logging_obj"),
user_api_key_dict=user_api_key_dict,
num_retries, retry_after = get_guardrail_retry_config(guardrail_to_apply)
await run_guardrail_with_retries(
coro_factory=lambda: endpoint_translation.process_output_streaming_response(
responses_so_far=responses_so_far,
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=request_data.get("litellm_logging_obj"),
user_api_key_dict=user_api_key_dict,
),
num_retries=num_retries,
retry_after=retry_after,
guardrail_name=guardrail_to_apply.guardrail_name,
)

View file

@ -469,6 +469,21 @@ class InMemoryGuardrailHandler:
else:
raise ValueError(f"Unsupported guardrail: {guardrail_type}")
# Inject guardrail retry config into callback so get_guardrail_retry_config can read it
if custom_guardrail_callback is not None:
if (
getattr(litellm_params, "num_retries", None) is not None
or getattr(litellm_params, "retry_after", None) is not None
):
opts = getattr(custom_guardrail_callback, "optional_params", None)
if opts is None:
setattr(custom_guardrail_callback, "optional_params", {})
opts = getattr(custom_guardrail_callback, "optional_params")
if getattr(litellm_params, "num_retries", None) is not None:
opts["num_retries"] = litellm_params.num_retries
if getattr(litellm_params, "retry_after", None) is not None:
opts["retry_after"] = litellm_params.retry_after
parsed_guardrail = Guardrail(
guardrail_id=guardrail.get("guardrail_id"),
guardrail_name=guardrail["guardrail_name"],

View file

@ -0,0 +1,161 @@
"""
Guardrail retry logic: mirrors router-level model retries (408, 409, 429, 5xx, no-status).
Never retries ModifyResponseException, ContentPolicyViolationError, NotFoundError.
"""
import asyncio
from typing import Any, Callable, Optional, TypeVar
import litellm
from litellm._logging import verbose_proxy_logger
T = TypeVar("T")
# Default retry config when not set on guardrail
DEFAULT_GUARDRAIL_NUM_RETRIES = 2
DEFAULT_GUARDRAIL_RETRY_AFTER = 0.0
def should_retry_guardrail_error(error: Exception) -> bool:
"""
Decide if a guardrail error is retriable.
- Never retry: ModifyResponseException, ContentPolicyViolationError, NotFoundError.
- Retry: 408, 409, 429, 5xx (via litellm._should_retry), and errors without status (e.g. network).
"""
from litellm.integrations.custom_guardrail import ModifyResponseException
if isinstance(error, ModifyResponseException):
return False
try:
import litellm.exceptions as litellm_exceptions
if isinstance(error, litellm_exceptions.ContentPolicyViolationError):
return False
if isinstance(error, litellm_exceptions.NotFoundError):
return False
except Exception:
pass
status_code: Optional[int] = getattr(error, "status_code", None)
if status_code is not None:
return litellm._should_retry(status_code)
# No status code (e.g. network error, timeout) -> retry
return True
def _time_to_sleep_before_guardrail_retry(
remaining_retries: int,
num_retries: int,
retry_after: float,
response_headers: Optional[Any] = None,
) -> float:
"""
Compute sleep before next guardrail retry using litellm._calculate_retry_after.
retry_after (config) is used as min_timeout for backoff.
"""
min_timeout = int(max(0.0, retry_after))
return litellm._calculate_retry_after(
remaining_retries=remaining_retries,
max_retries=num_retries,
response_headers=response_headers,
min_timeout=min_timeout,
)
async def run_guardrail_with_retries(
coro_factory: Callable[[], Any],
num_retries: int,
retry_after: float,
guardrail_name: str,
) -> Any:
"""
Run an async guardrail "call" with retries.
coro_factory: callable that returns a new coroutine for each attempt (so each attempt is fresh).
num_retries: max number of attempts (e.g. 2 means try once, then up to 2 retries = 3 total).
retry_after: minimum seconds to wait before retry (used in backoff).
guardrail_name: for logging.
Uses same semantics as router retry loop: retries on 408/409/429/5xx and no-status errors.
"""
if num_retries <= 0:
coro = coro_factory()
return await coro
last_error: Optional[Exception] = None
attempt = 0
remaining = num_retries
while True:
try:
coro = coro_factory()
return await coro
except Exception as e:
last_error = e
if not should_retry_guardrail_error(e):
raise
if remaining <= 0:
raise
response_headers = getattr(e, "response", None)
if response_headers is not None and hasattr(response_headers, "headers"):
response_headers = response_headers.headers
sleep_seconds = _time_to_sleep_before_guardrail_retry(
remaining_retries=remaining,
num_retries=num_retries,
retry_after=retry_after,
response_headers=response_headers,
)
verbose_proxy_logger.warning(
"Guardrail %s attempt %s failed (retriable): %s. Retrying in %.2fs (%s attempts left).",
guardrail_name,
attempt + 1,
type(e).__name__,
sleep_seconds,
remaining,
)
await asyncio.sleep(sleep_seconds)
remaining -= 1
attempt += 1
if last_error is not None:
raise last_error
def get_guardrail_retry_config(guardrail_to_apply: Any) -> tuple[int, float]:
"""
Read num_retries and retry_after from the guardrail instance.
Looks at optional_params and guardrail_config. Returns defaults when not set.
"""
num_retries = DEFAULT_GUARDRAIL_NUM_RETRIES
retry_after = DEFAULT_GUARDRAIL_RETRY_AFTER
opts = getattr(guardrail_to_apply, "optional_params", None)
if isinstance(opts, dict):
if "num_retries" in opts and opts["num_retries"] is not None:
num_retries = int(opts["num_retries"])
if "retry_after" in opts and opts["retry_after"] is not None:
retry_after = float(opts["retry_after"])
guardrail_config = getattr(guardrail_to_apply, "guardrail_config", None)
if isinstance(guardrail_config, dict):
if not (isinstance(opts, dict) and "num_retries" in opts) and "num_retries" in guardrail_config and guardrail_config["num_retries"] is not None:
num_retries = int(guardrail_config["num_retries"])
if not (isinstance(opts, dict) and "retry_after" in opts) and "retry_after" in guardrail_config and guardrail_config["retry_after"] is not None:
retry_after = float(guardrail_config["retry_after"])
# Proxy-level defaults from router_settings (when guardrail does not set its own)
try:
from litellm.proxy.proxy_server import llm_router
if llm_router is not None:
s = llm_router.get_settings()
if s is not None:
if num_retries == DEFAULT_GUARDRAIL_NUM_RETRIES and s.get("guardrail_num_retries") is not None:
num_retries = int(s["guardrail_num_retries"])
if retry_after == DEFAULT_GUARDRAIL_RETRY_AFTER and s.get("guardrail_retry_after") is not None:
retry_after = float(s["guardrail_retry_after"])
except Exception:
pass
return num_retries, retry_after

View file

@ -100,6 +100,10 @@ from litellm.proxy.db.create_views import (
from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter
from litellm.proxy.db.log_db_metrics import log_db_metrics
from litellm.proxy.db.prisma_client import PrismaWrapper
from litellm.proxy.guardrails.guardrail_retries import (
get_guardrail_retry_config,
run_guardrail_with_retries,
)
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
UnifiedLLMGuardrails,
)
@ -871,27 +875,71 @@ class ProxyLogging:
target = unified_guardrail if use_unified else callback
if hook_type == "pre_call":
return await target.async_pre_call_hook(
user_api_key_dict=user_api_key_dict, # type: ignore
cache=self.call_details["user_api_key_cache"],
data=data,
call_type=call_type,
)
elif hook_type == "during_call":
return await target.async_moderation_hook(
data=data,
user_api_key_dict=user_api_key_dict, # type: ignore
call_type=call_type,
)
elif hook_type == "post_call":
return await target.async_post_call_success_hook(
user_api_key_dict=user_api_key_dict, # type: ignore
data=data,
response=response, # type: ignore
)
if use_unified:
# Retries are handled inside unified_guardrail
if hook_type == "pre_call":
return await target.async_pre_call_hook(
user_api_key_dict=user_api_key_dict, # type: ignore
cache=self.call_details["user_api_key_cache"],
data=data,
call_type=call_type,
)
elif hook_type == "during_call":
return await target.async_moderation_hook(
data=data,
user_api_key_dict=user_api_key_dict, # type: ignore
call_type=call_type,
)
elif hook_type == "post_call":
return await target.async_post_call_success_hook(
user_api_key_dict=user_api_key_dict, # type: ignore
data=data,
response=response, # type: ignore
)
else:
raise ValueError(f"Unknown hook_type: {hook_type}")
else:
raise ValueError(f"Unknown hook_type: {hook_type}")
# Direct path: wrap with guardrail retries
num_retries, retry_after = get_guardrail_retry_config(callback)
guardrail_name = getattr(
callback, "guardrail_name", type(callback).__name__
)
if hook_type == "pre_call":
return await run_guardrail_with_retries(
coro_factory=lambda: target.async_pre_call_hook(
user_api_key_dict=user_api_key_dict, # type: ignore
cache=self.call_details["user_api_key_cache"],
data=data,
call_type=call_type,
),
num_retries=num_retries,
retry_after=retry_after,
guardrail_name=guardrail_name,
)
elif hook_type == "during_call":
return await run_guardrail_with_retries(
coro_factory=lambda: target.async_moderation_hook(
data=data,
user_api_key_dict=user_api_key_dict, # type: ignore
call_type=call_type,
),
num_retries=num_retries,
retry_after=retry_after,
guardrail_name=guardrail_name,
)
elif hook_type == "post_call":
return await run_guardrail_with_retries(
coro_factory=lambda: target.async_post_call_success_hook(
user_api_key_dict=user_api_key_dict, # type: ignore
data=data,
response=response, # type: ignore
),
num_retries=num_retries,
retry_after=retry_after,
guardrail_name=guardrail_name,
)
else:
raise ValueError(f"Unknown hook_type: {hook_type}")
async def _execute_guardrail_with_load_balancing(
self,
@ -1325,10 +1373,20 @@ class ProxyLogging:
call_type=call_type,
)
else:
guardrail_task = callback.async_moderation_hook(
data=data,
user_api_key_dict=user_api_key_auth_dict, # type: ignore
call_type=call_type, # type: ignore
_cb = callback
num_retries, retry_after = get_guardrail_retry_config(_cb)
guardrail_name = getattr(
_cb, "guardrail_name", type(_cb).__name__
)
guardrail_task = run_guardrail_with_retries(
coro_factory=lambda _c=_cb: _c.async_moderation_hook(
data=data,
user_api_key_dict=user_api_key_auth_dict, # type: ignore
call_type=call_type, # type: ignore
),
num_retries=num_retries,
retry_after=retry_after,
guardrail_name=guardrail_name,
)
guardrail_tasks.append(guardrail_task)
@ -1786,10 +1844,23 @@ class ProxyLogging:
)
)
else:
guardrail_response = await callback.async_post_call_success_hook(
user_api_key_dict=user_api_key_dict,
data=data,
response=response,
_guardrail_cb = callback
num_retries, retry_after = get_guardrail_retry_config(
_guardrail_cb
)
guardrail_name = getattr(
_guardrail_cb, "guardrail_name",
type(_guardrail_cb).__name__,
)
guardrail_response = await run_guardrail_with_retries(
coro_factory=lambda: _guardrail_cb.async_post_call_success_hook(
user_api_key_dict=user_api_key_dict,
data=data,
response=response,
),
num_retries=num_retries,
retry_after=retry_after,
guardrail_name=guardrail_name,
)
if guardrail_response is not None:

View file

@ -503,6 +503,10 @@ class Router:
self.retry_after = retry_after
self.routing_strategy = routing_strategy
# Proxy-level defaults for guardrail retries (used when guardrail config does not set num_retries/retry_after)
self.guardrail_num_retries = None
self.guardrail_retry_after = None
## SETTING FALLBACKS ##
### validate if it's set + in correct format
_fallbacks = fallbacks or litellm.fallbacks
@ -7586,6 +7590,8 @@ class Router:
"model_group_retry_policy",
"retry_policy",
"model_group_alias",
"guardrail_num_retries",
"guardrail_retry_after",
]
for var in vars_to_include:
@ -7616,6 +7622,8 @@ class Router:
"context_window_fallbacks",
"model_group_retry_policy",
"model_group_alias",
"guardrail_num_retries",
"guardrail_retry_after",
]
_int_settings = [
@ -7624,6 +7632,7 @@ class Router:
"retry_after",
"allowed_fails",
"cooldown_time",
"guardrail_num_retries",
]
_existing_router_settings = self.get_settings()

View file

@ -657,6 +657,16 @@ class BaseLitellmParams(
description="Python-like code containing the apply_guardrail function for custom guardrail logic",
)
# Guardrail retry config (mirrors router num_retries / retry_after)
num_retries: Optional[int] = Field(
default=None,
description="Number of retries for failed guardrail calls (e.g. default 2 when not set; 0 to disable)",
)
retry_after: Optional[float] = Field(
default=None,
description="Minimum seconds to wait before retrying a failed guardrail call (used in backoff)",
)
model_config = ConfigDict(extra="allow", protected_namespaces=())

View file

@ -189,6 +189,22 @@ ROUTER_SETTINGS_FIELDS: List[RouterSettingsField] = [
field_default=0,
ui_field_name="Retry After",
),
RouterSettingsField(
field_name="guardrail_num_retries",
field_type="Integer",
field_value=None,
field_description="Default number of retries for failed guardrail calls (used when guardrail config does not set num_retries)",
field_default=None,
ui_field_name="Guardrail Num Retries",
),
RouterSettingsField(
field_name="guardrail_retry_after",
field_type="Float",
field_value=None,
field_description="Default minimum seconds to wait before retrying a failed guardrail call",
field_default=None,
ui_field_name="Guardrail Retry After",
),
RouterSettingsField(
field_name="retry_policy",
field_type="Dictionary",

View file

@ -0,0 +1,148 @@
"""Tests for guardrail retry logic (guardrail_retries.py)."""
import pytest
from unittest.mock import AsyncMock, patch
import litellm
from litellm.integrations.custom_guardrail import ModifyResponseException
from litellm.proxy.guardrails.guardrail_retries import (
DEFAULT_GUARDRAIL_NUM_RETRIES,
DEFAULT_GUARDRAIL_RETRY_AFTER,
get_guardrail_retry_config,
run_guardrail_with_retries,
should_retry_guardrail_error,
)
class TestShouldRetryGuardrailError:
"""Test which errors are considered retriable."""
def test_retries_on_429(self):
e = Exception("rate limit")
e.status_code = 429
assert should_retry_guardrail_error(e) is True
def test_retries_on_500(self):
e = Exception("server error")
e.status_code = 500
assert should_retry_guardrail_error(e) is True
def test_no_retry_on_404(self):
e = Exception("not found")
e.status_code = 404
assert should_retry_guardrail_error(e) is False
def test_no_retry_on_modify_response_exception(self):
e = ModifyResponseException(
message="blocked",
model="gpt-4",
request_data={},
)
assert should_retry_guardrail_error(e) is False
def test_no_retry_on_content_policy_violation(self):
try:
e = litellm.ContentPolicyViolationError(
message="policy", model="gpt-4", llm_provider="openai"
)
assert should_retry_guardrail_error(e) is False
except (AttributeError, TypeError):
pytest.skip("ContentPolicyViolationError not available")
class TestGetGuardrailRetryConfig:
"""Test reading num_retries and retry_after from guardrail."""
def test_defaults(self):
guardrail = type("G", (), {})()
num_retries, retry_after = get_guardrail_retry_config(guardrail)
assert num_retries == DEFAULT_GUARDRAIL_NUM_RETRIES
assert retry_after == DEFAULT_GUARDRAIL_RETRY_AFTER
def test_from_optional_params(self):
guardrail = type("G", (), {"optional_params": {"num_retries": 5, "retry_after": 2.0}})()
num_retries, retry_after = get_guardrail_retry_config(guardrail)
assert num_retries == 5
assert retry_after == 2.0
def test_from_guardrail_config(self):
guardrail = type("G", (), {"guardrail_config": {"num_retries": 3, "retry_after": 1}})()
num_retries, retry_after = get_guardrail_retry_config(guardrail)
assert num_retries == 3
assert retry_after == 1
class TestRunGuardrailWithRetries:
"""Test run_guardrail_with_retries behavior."""
@pytest.mark.asyncio
async def test_succeeds_first_try(self):
async def ok_coro():
return {"done": True}
result = await run_guardrail_with_retries(
coro_factory=lambda: ok_coro(),
num_retries=2,
retry_after=0,
guardrail_name="test",
)
assert result == {"done": True}
@pytest.mark.asyncio
async def test_succeeds_on_second_attempt(self):
call_count = 0
async def fail_once():
nonlocal call_count
call_count += 1
if call_count == 1:
e = Exception("retriable")
e.status_code = 429
raise e
return {"done": True}
with patch("litellm.proxy.guardrails.guardrail_retries.asyncio.sleep", new_callable=AsyncMock):
result = await run_guardrail_with_retries(
coro_factory=lambda: fail_once(),
num_retries=2,
retry_after=0,
guardrail_name="test",
)
assert result == {"done": True}
assert call_count == 2
@pytest.mark.asyncio
async def test_raises_after_exhausting_retries(self):
async def always_fail():
e = Exception("rate limit")
e.status_code = 429
raise e
with patch("litellm.proxy.guardrails.guardrail_retries.asyncio.sleep", new_callable=AsyncMock):
with pytest.raises(Exception, match="rate limit"):
await run_guardrail_with_retries(
coro_factory=lambda: always_fail(),
num_retries=2,
retry_after=0,
guardrail_name="test",
)
@pytest.mark.asyncio
async def test_no_retry_when_num_retries_zero(self):
call_count = 0
async def fail_once():
nonlocal call_count
call_count += 1
raise Exception("boom")
with patch("litellm.proxy.guardrails.guardrail_retries.asyncio.sleep", new_callable=AsyncMock):
with pytest.raises(Exception, match="boom"):
await run_guardrail_with_retries(
coro_factory=lambda: fail_once(),
num_retries=0,
retry_after=0,
guardrail_name="test_guardrail",
)
assert call_count == 1