diff --git a/litellm/router.py b/litellm/router.py index ddf2d5f4c68..ded1f75fd38 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -4953,22 +4953,30 @@ class Router: """ Retry Logic """ - # For guardrail calls, use guardrail_list so should_retry_this_error allows retries - _model = kwargs.get("model") or "" - _guardrail_deployments = [ - g for g in self.guardrail_list if g.get("guardrail_name") == _model - ] - if kwargs.get("selected_guardrail") is not None or _guardrail_deployments: - _healthy_deployments = _guardrail_deployments + # For guardrail calls, use guardrail_list so should_retry_this_error allows retries. + # When selected_guardrail is in kwargs (normal path), avoid scanning guardrail_list. + _selected_guardrail = kwargs.get("selected_guardrail") + if _selected_guardrail is not None: + _healthy_deployments = [_selected_guardrail] _all_deployments = _healthy_deployments else: - ( - _healthy_deployments, - _all_deployments, - ) = await self._async_get_healthy_deployments( - model=kwargs.get("model") or "", - parent_otel_span=parent_otel_span, - ) + _model = kwargs.get("model") or "" + _guardrail_deployments = [ + g + for g in self.guardrail_list + if g.get("guardrail_name") == _model + ] + if _guardrail_deployments: + _healthy_deployments = _guardrail_deployments + _all_deployments = _healthy_deployments + else: + ( + _healthy_deployments, + _all_deployments, + ) = await self._async_get_healthy_deployments( + model=_model, + parent_otel_span=parent_otel_span, + ) # Check retry policy FIRST, before should_retry_this_error # This allows retry policies to override the healthy deployments check @@ -5043,30 +5051,21 @@ class Router: kwargs = self.log_retry(kwargs=kwargs, e=e) remaining_retries = num_retries - current_attempt - 1 _model = kwargs.get("model") # type: ignore - _is_guardrail = ( - kwargs.get("selected_guardrail") is not None - or any( - g.get("guardrail_name") == _model - for g in self.guardrail_list - ) - ) - if _is_guardrail and _model: - _healthy_deployments = [ - g - for g in self.guardrail_list - if g.get("guardrail_name") == _model - ] + _sel = kwargs.get("selected_guardrail") + if _sel is not None: + _healthy_deployments = [_sel] _all_deployments = _healthy_deployments elif _model is not None: ( _healthy_deployments, - _, + _all_deployments, ) = await self._async_get_healthy_deployments( model=_model, parent_otel_span=parent_otel_span, ) else: _healthy_deployments = [] + _all_deployments = [] _timeout = self._time_to_sleep_before_retry( e=original_exception, remaining_retries=remaining_retries, diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 3ac04deacbd..fe85b7edbf6 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -1872,7 +1872,7 @@ async def test_aguardrail_retry_succeeds_after_retries(): raise ValueError("guardrail API temporarily failed") return {"result": "success", "attempt": call_count} - # Skip sleep so test stays fast; guardrail retries use guardrail_list for healthy_deployments (no patch needed) + # Skip sleep so test stays fast with patch.object(router, "_time_to_sleep_before_retry", return_value=0): result = await router.aguardrail( guardrail_name="flaky-guardrail", @@ -1917,7 +1917,7 @@ async def test_aguardrail_retry_exhausted_raises(): call_count += 1 raise ValueError("guardrail API always fails") - # Skip sleep so test stays fast; guardrail retries use guardrail_list for healthy_deployments (no patch needed) + # Skip sleep so test stays fast with patch.object(router, "_time_to_sleep_before_retry", return_value=0): with pytest.raises(ValueError, match="guardrail API always fails"): await router.aguardrail(