mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
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:
commit
80b9ed4f2c
10 changed files with 106 additions and 24 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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue