mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-15 23:31:29 +00:00
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:
parent
6764ab2673
commit
aaf924693a
8 changed files with 85 additions and 30 deletions
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue