Merge pull request #41191 from BerriAI/litellm_router_test_cap_resets_per_fallback_hop

fix(router): count num_retries_per_request across fallback hops
This commit is contained in:
Mateo Wang 2026-09-15 01:13:41 -07:00 committed by GitHub
commit 80b9ed4f2c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 106 additions and 24 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

@ -336,7 +336,7 @@ _CLIENT_PRICING_METADATA_FIELDS: Final = frozenset({"model_info", "standard_logg
# and read by spend logs as fact; a client value has no legitimate meaning and no
# key or team setting keeps it, so the strip is never gated.
_ROUTER_RESERVED_METADATA_FIELDS: Final = frozenset(
{"attempted_fallbacks", "original_model_group", CLIENT_OUTPUT_CEILING_METADATA_KEY}
{"attempted_fallbacks", "original_model_group", "request_retry_count", CLIENT_OUTPUT_CEILING_METADATA_KEY}
)
_ALLOW_CLIENT_PRICING_OVERRIDE_METADATA_KEY: Final = "allow_client_pricing_override"

View file

@ -8377,7 +8377,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
@ -8401,7 +8402,10 @@ class Router:
else ()
)
breadcrumbs: Final = (*kept_breadcrumbs, attempt_record)
earlier: Final = request_metadata.get("request_retry_count")
request_retry_count: Final = (earlier if type(earlier) is int and 0 <= earlier 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

@ -11,7 +11,7 @@ import litellm
from unittest.mock import patch, MagicMock, AsyncMock
from create_mock_standard_logging_payload import create_standard_logging_payload
from litellm.types.utils import StandardLoggingPayload
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo
from litellm.types.router import Deployment, DeploymentTypedDict, LiteLLM_Params, ModelInfo
from litellm.constants import DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS
@ -630,10 +630,12 @@ 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"""
def test_log_retry(model_list: list[DeploymentTypedDict], metadata_key: str) -> None:
"""log_retry appends one flat record per failed attempt, copies neither the request kwargs nor the
request metadata into it, counts every failed attempt of the request independently of the
per-hop attempted_retries, and never trusts a negative count planted before the first failure"""
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,10 @@ 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
planted_kwargs = {"model": "gpt-3.5-turbo", metadata_key: {"request_retry_count": -100}}
assert router.log_retry(kwargs=planted_kwargs, e=rate_limit_error)[metadata_key]["request_retry_count"] == 1
def test_update_usage(model_list):

View file

@ -7352,13 +7352,14 @@ def _reserved_stamp_key(key_metadata: dict | None = None) -> UserAPIKeyAuth:
_PLANTED_STAMPS = {
"attempted_fallbacks": 99,
"original_model_group": "spoofed-group",
"request_retry_count": -100,
"_client_output_ceiling": {"api_base": "https://attacker.example"},
"client_key": "client_value",
}
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_strips_router_reserved_stamps_from_both_buckets():
async def test_add_litellm_data_to_request_strips_router_reserved_stamps_from_both_buckets() -> None:
"""attempted_fallbacks and original_model_group are router-written facts the spend row
reads back; a client planting them in either bucket is dropped at the boundary so the
router never sees a reserved key it did not write."""
@ -7384,11 +7385,12 @@ async def test_add_litellm_data_to_request_strips_router_reserved_stamps_from_bo
assert "attempted_fallbacks" not in updated["metadata"]
assert "original_model_group" not in updated["metadata"]
assert "_client_output_ceiling" not in updated["metadata"]
assert "request_retry_count" not in updated["metadata"]
assert updated["metadata"]["client_key"] == "client_value"
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_strips_router_reserved_stamps_from_json_string_litellm_metadata():
async def test_add_litellm_data_to_request_strips_router_reserved_stamps_from_json_string_litellm_metadata() -> None:
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
data = {
@ -7409,11 +7411,12 @@ async def test_add_litellm_data_to_request_strips_router_reserved_stamps_from_js
assert "litellm_metadata" not in updated
assert "attempted_fallbacks" not in updated["metadata"]
assert "original_model_group" not in updated["metadata"]
assert "request_retry_count" not in updated["metadata"]
assert updated["metadata"]["client_key"] == "client_value"
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_strips_router_reserved_stamps_despite_pricing_override_opt_in():
async def test_add_litellm_data_to_request_strips_router_reserved_stamps_despite_pricing_override_opt_in() -> None:
"""The pricing strip is gated on allow_client_pricing_override; the reserved-stamp strip
is not, because no key or team setting makes a client-written fallback count valid."""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
@ -7437,6 +7440,7 @@ async def test_add_litellm_data_to_request_strips_router_reserved_stamps_despite
assert updated["metadata"]["model_info"] == {"input_cost_per_token": 0.0}
assert "attempted_fallbacks" not in updated["metadata"]
assert "original_model_group" not in updated["metadata"]
assert "request_retry_count" not in updated["metadata"]
@pytest.mark.asyncio

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,66 @@ async def test_num_retries_per_request_stops_retries_at_caps_above_four(monkeypa
]
def _failing_group_with_healthy_fallback_router(num_retries: int) -> litellm.Router:
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, planted_count, hop_refused",
[(2, None, True), (4, None, False), (2, -100, True)],
ids=["cap-spent-before-the-hop", "cap-not-reached-by-the-hop", "planted-negative-count-does-not-lift-the-cap"],
)
async def test_num_retries_per_request_counts_retries_across_fallback_hops(
monkeypatch: pytest.MonkeyPatch, cap: int, planted_count: int | None, hop_refused: bool
) -> None:
"""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. A caller who plants a negative count
in the request metadata must not push the cap further away either."""
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:
metadata = {} if planted_count is None else {"request_retry_count": planted_count}
request = router.acompletion(
model="broken-group", messages=[{"role": "user", "content": "hi"}], metadata=metadata
)
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,9 @@ 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: pytest.MonkeyPatch, metadata_key: str, cap: int, metadata: object, refused: bool
) -> None:
monkeypatch.setattr(litellm, "num_retries_per_request", cap)
kwargs: Final = _capped_completion_kwargs(metadata_key, metadata)
if refused:
@ -4111,7 +4114,9 @@ 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: pytest.MonkeyPatch, metadata_key: str, cap: int, metadata: object, refused: bool
) -> None:
monkeypatch.setattr(litellm, "num_retries_per_request", cap)
kwargs: Final = _capped_completion_kwargs(metadata_key, metadata)
if refused:

View file

@ -806,7 +806,7 @@ async def test_shared_call_limits_still_reject_before_reading_ocr_file(
monkeypatch.setattr(litellm, "_current_cost", 2)
monkeypatch.setattr(litellm, "num_retries_per_request", 1 if limit == "retries" else None)
expected: Final = litellm.BudgetExceededError if limit == "budget" else RuntimeError
arguments: Final = {"document": {"type": "file", "file": File()}, "metadata": {"attempted_retries": 1}}
arguments: Final = {"document": {"type": "file", "file": File()}, "metadata": {"request_retry_count": 1}}
with pytest.raises(expected, match=r"Budget has been exceeded|Max retries per request hit"):
await call_aocr(ocr_server, **arguments) if asynchronous else call_ocr(ocr_server, **arguments)
assert reads == []