mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
added guardrail retry logic
This commit is contained in:
parent
0649720f79
commit
f2370bd3ba
8 changed files with 517 additions and 51 deletions
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
161
litellm/proxy/guardrails/guardrail_retries.py
Normal file
161
litellm/proxy/guardrails/guardrail_retries.py
Normal 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
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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=())
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
148
tests/test_litellm/proxy/guardrails/test_guardrail_retries.py
Normal file
148
tests/test_litellm/proxy/guardrails/test_guardrail_retries.py
Normal 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
|
||||
Loading…
Add table
Reference in a new issue