fix(router): count num_retries_per_request across fallback hops

num_retries_per_request has always capped the retries of one request with its fallback hops included. #40930 started reading the per-hop attempted_retries counter instead, and every fallback hop restarts that counter at zero, so a request could spend a fresh retry budget on each hop and the legacy fallback cap test started seeing the hop run.

Router.log_retry now also keeps request_retry_count on the request metadata, incremented on every retry and fallback hop and never truncated the way previous_models is, and max_retries_per_request_hit reads that count. The flat retry records, the litellm_metadata coverage and caps above four from #40930 stay as they are, and the legacy test goes back to its previous_models == 0 assertion.
This commit is contained in:
mateo-berri 2026-09-14 23:13:50 -07:00
parent 6764ab2673
commit aaf924693a
8 changed files with 85 additions and 30 deletions

View file

@ -538,7 +538,7 @@ context_window_fallbacks: Optional[List] = None
content_policy_fallbacks: Optional[List] = None
allowed_fails: int = 3
allow_dynamic_callback_disabling: bool = True
num_retries_per_request: Optional[int] = None # cap on Router retries of one model group; resets per fallback hop
num_retries_per_request: Optional[int] = None # for the request overall (incl. fallbacks + model retries)
####### SECRET MANAGERS #####################
secret_manager_client: Optional[Any] = (
None # list of instantiated key management clients - e.g. azure kv, infisical, etc.

View file

@ -309,8 +309,8 @@ def max_retries_per_request_hit(kwargs: Mapping[str, object], num_retries_per_re
metadata: Final = kwargs.get(get_metadata_variable_name_from_kwargs(kwargs))
if not isinstance(metadata, Mapping):
return False
attempted_retries: Final = metadata.get("attempted_retries")
return type(attempted_retries) is int and 0 < attempted_retries and num_retries_per_request <= attempted_retries
retry_count: Final = metadata.get("request_retry_count")
return type(retry_count) is int and 0 < retry_count and num_retries_per_request <= retry_count
def get_or_create_metadata_bucket(

View file

@ -8378,7 +8378,8 @@ class Router:
def log_retry(self, kwargs: dict, e: Exception) -> dict:
"""
When a retry or fallback happens, record which model group, deployment and attempt just failed and why
When a retry or fallback happens, record which model group, deployment and attempt just failed and why,
and count it toward the request-wide num_retries_per_request cap
"""
from litellm.types.router import RetryAttemptRecord
@ -8402,7 +8403,10 @@ class Router:
else ()
)
breadcrumbs: Final = (*kept_breadcrumbs, attempt_record)
earlier_retry_count: Final = request_metadata.get("request_retry_count")
request_retry_count: Final = (earlier_retry_count if type(earlier_retry_count) is int else 0) + 1
kwargs[_metadata_var]["previous_models"] = breadcrumbs # rebind-ok: the logging object already holds this dict
kwargs[_metadata_var]["request_retry_count"] = request_retry_count # rebind-ok: same dict, read by the cap
return kwargs
def _update_usage(self, deployment_id: str, parent_otel_span: Span | None) -> int:

View file

@ -5,7 +5,6 @@ import asyncio
import os
import time
import traceback
from typing import Final
import pytest
@ -14,7 +13,6 @@ from unittest.mock import AsyncMock, MagicMock, patch
import litellm
from litellm import Router
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.router import RetryAttemptRecord
from tests.fake_openai_endpoint import FAKE_OPENAI_API_BASE
@ -23,15 +21,15 @@ class MyCustomHandler(CustomLogger):
success: bool = False
failure: bool = False
previous_models: int = 0
previous_model_records: tuple[RetryAttemptRecord, ...] = ()
def log_pre_api_call(self, model, messages, kwargs):
print(f"Pre-API Call")
print(
f"previous_models: {kwargs['litellm_params']['metadata'].get('previous_models', None)}"
)
self.previous_model_records = tuple(kwargs["litellm_params"]["metadata"].get("previous_models", ()))
self.previous_models = len(self.previous_model_records)
self.previous_models = len(
kwargs["litellm_params"]["metadata"].get("previous_models", [])
) # {"previous_models": [{"model": litellm_model_name, "exception_type": AuthenticationError, "exception_string": <complete_traceback>}]}
print(f"self.previous_models: {self.previous_models}")
def log_post_api_call(self, kwargs, response_obj, start_time, end_time):
@ -720,14 +718,7 @@ async def test_async_fallbacks_max_retries_per_request():
await asyncio.sleep(
0.05
) # allow a delay as success_callbacks are on a separate thread
records: Final = customHandler.previous_model_records
assert customHandler.previous_models == len(records)
assert records
assert {record["model_group"] for record in records} == {"azure/gpt-3.5-turbo"}
assert next(record["exception_type"] for record in records if record["attempted_retries"] == 0) == "AuthenticationError"
refused_retries: Final = tuple(record for record in records if record["attempted_retries"])
assert refused_retries
assert all("Max retries per request hit!" in record["exception_string"] for record in refused_retries)
assert customHandler.previous_models == 0 # 0 retries, 0 fallback
router.reset()
except litellm.Timeout as e:
pass

View file

@ -631,9 +631,11 @@ def test_deployment_callback_respects_cooldown_time(model_list):
@pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"])
def test_log_retry(model_list, metadata_key):
"""log_retry appends one flat record per failed attempt and copies neither the request kwargs nor
the request metadata into it"""
"""log_retry appends one flat record per failed attempt, copies neither the request kwargs nor the
request metadata into it, and counts every failed attempt of the request independently of the
per-hop attempted_retries"""
router = Router(model_list=model_list)
rate_limit_error = litellm.RateLimitError(message="slow down", llm_provider="openai", model="gpt-3.5-turbo")
new_kwargs = router.log_retry(
kwargs={
"model": "gpt-3.5-turbo",
@ -641,7 +643,7 @@ def test_log_retry(model_list, metadata_key):
"messages": [{"role": "user", "content": "hi"}],
metadata_key: {"model_info": {"id": "deployment-1"}, "attempted_retries": 2, "user_api_key": "sk-proxy"},
},
e=litellm.RateLimitError(message="slow down", llm_provider="openai", model="gpt-3.5-turbo"),
e=rate_limit_error,
)
assert json.loads(json.dumps(new_kwargs[metadata_key]["previous_models"])) == [
{
@ -652,6 +654,8 @@ def test_log_retry(model_list, metadata_key):
"attempted_retries": 2,
}
]
assert new_kwargs[metadata_key]["request_retry_count"] == 1
assert router.log_retry(kwargs=new_kwargs, e=rate_limit_error)[metadata_key]["request_retry_count"] == 2
def test_update_usage(model_list):

View file

@ -8,7 +8,7 @@ from litellm.rust_bridge.lifecycle import check_limits
@pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"])
@pytest.mark.parametrize(
"cap, attempted_retries, refused",
"cap, request_retry_count, refused",
[(5, 5, True), (5, 4, False), (0, 0, False), (0, 1, True)],
ids=[
"cap-above-four-reached",
@ -17,12 +17,15 @@ from litellm.rust_bridge.lifecycle import check_limits
"cap-of-zero-refuses-first-retry",
],
)
def test_check_limits_reads_attempted_retries(
monkeypatch: pytest.MonkeyPatch, metadata_key: str, cap: int, attempted_retries: int, refused: bool
def test_check_limits_reads_request_retry_count(
monkeypatch: pytest.MonkeyPatch, metadata_key: str, cap: int, request_retry_count: int, refused: bool
) -> None:
monkeypatch.setattr(litellm, "num_retries_per_request", cap)
monkeypatch.setattr(litellm, "max_budget", None)
kwargs: Final = {"model": "mistral/mistral-ocr-latest", metadata_key: {"attempted_retries": attempted_retries}}
kwargs: Final = {
"model": "mistral/mistral-ocr-latest",
metadata_key: {"request_retry_count": request_retry_count},
}
if refused:
with pytest.raises(RuntimeError, match="Max retries per request hit!"):
check_limits(kwargs)

View file

@ -11066,6 +11066,58 @@ async def test_num_retries_per_request_stops_retries_at_caps_above_four(monkeypa
]
def _failing_group_with_healthy_fallback_router(num_retries):
return litellm.Router(
model_list=[
{
"model_name": "broken-group",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "sk-fake",
"mock_response": "litellm.InternalServerError",
},
},
{
"model_name": "healthy-group",
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "sk-fake", "mock_response": "ok"},
},
],
fallbacks=[{"broken-group": ["healthy-group"]}],
num_retries=num_retries,
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"cap, hop_refused", [(2, True), (4, False)], ids=["cap-spent-before-the-hop", "cap-not-reached-by-the-hop"]
)
async def test_num_retries_per_request_counts_retries_across_fallback_hops(monkeypatch, cap, hop_refused):
"""num_retries_per_request caps the retries of one request, fallback hops included. Each hop starts a
fresh per-hop attempted_retries at zero, so a cap read from that counter let every hop retry from zero
and a request could spend far more retries than the cap allows."""
monkeypatch.setattr(litellm, "num_retries_per_request", cap)
router = _failing_group_with_healthy_fallback_router(num_retries=1)
recorder = _FallbackAttemptRecorder()
litellm.callbacks.append(recorder)
try:
request = router.acompletion(model="broken-group", messages=[{"role": "user", "content": "hi"}])
if not hop_refused:
assert (await request).choices[0].message.content == "ok"
return
with pytest.raises(litellm.InternalServerError):
await request
finally:
litellm.callbacks.remove(recorder)
assert recorder.failed_targets == ["healthy-group"]
hop_refusals = [
record["attempted_retries"]
for record in recorder.breadcrumbs_per_target[0]
if record["model_group"] == "healthy-group" and "Max retries per request hit!" in record["exception_string"]
]
assert hop_refusals == [0, 1]
@pytest.mark.asyncio
async def test_fallback_traceback_stays_available_at_debug_level():
"""Dropping the stack from the ERROR line is only safe because the fallback path still

View file

@ -4077,10 +4077,11 @@ class TestMetadataNoneHandling:
_RETRY_CAP_CASES: Final = (
pytest.param(5, {"attempted_retries": 5}, True, id="cap-above-four-reached"),
pytest.param(5, {"attempted_retries": 4}, False, id="cap-above-four-not-reached"),
pytest.param(0, {"attempted_retries": 0}, False, id="first-attempt-passes-cap-of-zero"),
pytest.param(0, {"attempted_retries": 1}, True, id="cap-of-zero-refuses-first-retry"),
pytest.param(5, {"request_retry_count": 5}, True, id="cap-above-four-reached"),
pytest.param(5, {"request_retry_count": 4}, False, id="cap-above-four-not-reached"),
pytest.param(0, {"request_retry_count": 0}, False, id="first-attempt-passes-cap-of-zero"),
pytest.param(0, {"request_retry_count": 1}, True, id="cap-of-zero-refuses-first-retry"),
pytest.param(0, {"attempted_retries": 1}, False, id="per-hop-attempted-retries-is-not-the-cap"),
pytest.param(5, {"previous_models": ("a", "b", "c", "d", "e")}, False, id="breadcrumb-count-is-not-the-cap"),
pytest.param(5, None, False, id="metadata-none"),
)
@ -4098,7 +4099,7 @@ def _capped_completion_kwargs(metadata_key: str, metadata: object) -> dict[str,
@pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"])
@pytest.mark.parametrize("cap, metadata, refused", _RETRY_CAP_CASES)
def test_num_retries_per_request_reads_attempted_retries_sync(monkeypatch, metadata_key, cap, metadata, refused):
def test_num_retries_per_request_reads_request_retry_count_sync(monkeypatch, metadata_key, cap, metadata, refused):
monkeypatch.setattr(litellm, "num_retries_per_request", cap)
kwargs: Final = _capped_completion_kwargs(metadata_key, metadata)
if refused:
@ -4111,7 +4112,7 @@ def test_num_retries_per_request_reads_attempted_retries_sync(monkeypatch, metad
@pytest.mark.asyncio
@pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"])
@pytest.mark.parametrize("cap, metadata, refused", _RETRY_CAP_CASES)
async def test_num_retries_per_request_reads_attempted_retries_async(monkeypatch, metadata_key, cap, metadata, refused):
async def test_num_retries_per_request_reads_request_retry_count_async(monkeypatch, metadata_key, cap, metadata, refused):
monkeypatch.setattr(litellm, "num_retries_per_request", cap)
kwargs: Final = _capped_completion_kwargs(metadata_key, metadata)
if refused: