From f2370bd3ba678e62abe157576d5228c9985030fb Mon Sep 17 00:00:00 2001 From: shivam Date: Thu, 5 Feb 2026 14:25:53 -0800 Subject: [PATCH] added guardrail retry logic --- .../unified_guardrail/unified_guardrail.py | 82 ++++++--- .../proxy/guardrails/guardrail_registry.py | 15 ++ litellm/proxy/guardrails/guardrail_retries.py | 161 ++++++++++++++++++ litellm/proxy/utils.py | 127 +++++++++++--- litellm/router.py | 9 + litellm/types/guardrails.py | 10 ++ .../router_settings_endpoints.py | 16 ++ .../guardrails/test_guardrail_retries.py | 148 ++++++++++++++++ 8 files changed, 517 insertions(+), 51 deletions(-) create mode 100644 litellm/proxy/guardrails/guardrail_retries.py create mode 100644 tests/test_litellm/proxy/guardrails/test_guardrail_retries.py diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index f07f65d10f5..3cfa7386de2 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -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, ) diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index c3da6892209..3572ef566a9 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -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"], diff --git a/litellm/proxy/guardrails/guardrail_retries.py b/litellm/proxy/guardrails/guardrail_retries.py new file mode 100644 index 00000000000..253dee6178c --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_retries.py @@ -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 diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 6bbf0df74de..5d6d1cde8da 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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: diff --git a/litellm/router.py b/litellm/router.py index d01c8443dab..48ab6c76f20 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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() diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index 74ccb34ca6e..88b3f240dd6 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -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=()) diff --git a/litellm/types/management_endpoints/router_settings_endpoints.py b/litellm/types/management_endpoints/router_settings_endpoints.py index 5024fe39b37..96999e867cc 100644 --- a/litellm/types/management_endpoints/router_settings_endpoints.py +++ b/litellm/types/management_endpoints/router_settings_endpoints.py @@ -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", diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_retries.py b/tests/test_litellm/proxy/guardrails/test_guardrail_retries.py new file mode 100644 index 00000000000..7a7887c54a2 --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_retries.py @@ -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