mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
fix(router): honor per-deployment num_retries for Responses and Anthropic Messages
This commit is contained in:
parent
c274cf321c
commit
ab66db830c
2 changed files with 133 additions and 48 deletions
|
|
@ -4497,7 +4497,6 @@ class Router:
|
|||
"""
|
||||
|
||||
passthrough_on_no_deployment = kwargs.pop("passthrough_on_no_deployment", False)
|
||||
function_name = "_ageneric_api_call_with_fallbacks"
|
||||
try:
|
||||
parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs)
|
||||
try:
|
||||
|
|
@ -4513,65 +4512,73 @@ class Router:
|
|||
return await original_generic_function(model=model, **kwargs)
|
||||
raise e
|
||||
|
||||
self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs, function_name=function_name)
|
||||
|
||||
data = deployment["litellm_params"].copy()
|
||||
model_name = data["model"]
|
||||
self.total_calls[model_name] += 1
|
||||
|
||||
self._add_deployment_model_to_endpoint_for_llm_passthrough_route(
|
||||
kwargs=kwargs, model=model, model_name=model_name
|
||||
)
|
||||
|
||||
# Get custom_llm_provider from deployment params
|
||||
try:
|
||||
custom_llm_provider = data.get("custom_llm_provider")
|
||||
_, inferred_custom_llm_provider, _, _ = get_llm_provider(
|
||||
model=data["model"],
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
function_name = "_ageneric_api_call_with_fallbacks"
|
||||
self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs, function_name=function_name)
|
||||
|
||||
data = deployment["litellm_params"].copy()
|
||||
model_name = data["model"]
|
||||
self.total_calls[model_name] += 1
|
||||
|
||||
self._add_deployment_model_to_endpoint_for_llm_passthrough_route(
|
||||
kwargs=kwargs, model=model, model_name=model_name
|
||||
)
|
||||
custom_llm_provider = custom_llm_provider or inferred_custom_llm_provider
|
||||
except Exception:
|
||||
custom_llm_provider = None
|
||||
|
||||
response_kwargs = {
|
||||
**data,
|
||||
"caching": self.cache_responses,
|
||||
**kwargs,
|
||||
"model": model_name,
|
||||
}
|
||||
# Only set custom_llm_provider if it's not None
|
||||
if custom_llm_provider is not None:
|
||||
response_kwargs["custom_llm_provider"] = custom_llm_provider
|
||||
# Get custom_llm_provider from deployment params
|
||||
try:
|
||||
custom_llm_provider = data.get("custom_llm_provider")
|
||||
_, inferred_custom_llm_provider, _, _ = get_llm_provider(
|
||||
model=data["model"],
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
custom_llm_provider = custom_llm_provider or inferred_custom_llm_provider
|
||||
except Exception:
|
||||
custom_llm_provider = None
|
||||
|
||||
response = original_generic_function(**response_kwargs)
|
||||
response_kwargs = {
|
||||
**data,
|
||||
"caching": self.cache_responses,
|
||||
**kwargs,
|
||||
"model": model_name,
|
||||
}
|
||||
# Only set custom_llm_provider if it's not None
|
||||
if custom_llm_provider is not None:
|
||||
response_kwargs["custom_llm_provider"] = custom_llm_provider
|
||||
|
||||
rpm_semaphore = self._get_client(
|
||||
deployment=deployment,
|
||||
kwargs=kwargs,
|
||||
client_type="max_parallel_requests",
|
||||
)
|
||||
response = original_generic_function(**response_kwargs)
|
||||
|
||||
if rpm_semaphore is not None and isinstance(rpm_semaphore, asyncio.Semaphore):
|
||||
async with rpm_semaphore:
|
||||
"""
|
||||
- Check rpm limits before making the call
|
||||
- If allowed, increment the rpm limit (allows global value to be updated, concurrency-safe)
|
||||
"""
|
||||
rpm_semaphore = self._get_client(
|
||||
deployment=deployment,
|
||||
kwargs=kwargs,
|
||||
client_type="max_parallel_requests",
|
||||
)
|
||||
|
||||
if rpm_semaphore is not None and isinstance(rpm_semaphore, asyncio.Semaphore):
|
||||
async with rpm_semaphore:
|
||||
"""
|
||||
- Check rpm limits before making the call
|
||||
- If allowed, increment the rpm limit (allows global value to be updated, concurrency-safe)
|
||||
"""
|
||||
await self.async_routing_strategy_pre_call_checks(
|
||||
deployment=deployment, parent_otel_span=parent_otel_span
|
||||
)
|
||||
response = await response # type: ignore
|
||||
else:
|
||||
await self.async_routing_strategy_pre_call_checks(
|
||||
deployment=deployment, parent_otel_span=parent_otel_span
|
||||
)
|
||||
response = await response # type: ignore
|
||||
else:
|
||||
await self.async_routing_strategy_pre_call_checks(
|
||||
deployment=deployment, parent_otel_span=parent_otel_span
|
||||
|
||||
self.success_calls[model_name] += 1
|
||||
verbose_router_logger.info(
|
||||
f"ageneric_api_call_with_fallbacks(model={model_name})\033[32m 200 OK\033[0m"
|
||||
)
|
||||
response = await response # type: ignore
|
||||
|
||||
self.success_calls[model_name] += 1
|
||||
verbose_router_logger.info(f"ageneric_api_call_with_fallbacks(model={model_name})\033[32m 200 OK\033[0m")
|
||||
|
||||
return response
|
||||
return response
|
||||
except Exception as e:
|
||||
self._set_deployment_num_retries_on_exception(e, deployment)
|
||||
self._set_failed_deployment_id_on_exception(e, deployment)
|
||||
raise e
|
||||
except Exception as e:
|
||||
verbose_router_logger.info(
|
||||
f"ageneric_api_call_with_fallbacks(model={model})\033[31m Exception {str(e)}\033[0m"
|
||||
|
|
|
|||
|
|
@ -575,3 +575,81 @@ class TestRequestNumRetriesBeatsGlobal:
|
|||
model="mock", messages=[{"role": "user", "content": "hi"}]
|
||||
)
|
||||
assert counter.attempts == 3
|
||||
|
||||
|
||||
class TestGenericApiCallDeploymentNumRetries:
|
||||
"""
|
||||
The Responses API and Anthropic Messages routes go through
|
||||
``_ageneric_api_call_with_fallbacks`` rather than ``_acompletion``, and that path used
|
||||
to swallow the selected deployment's ``litellm_params.num_retries``: the raised
|
||||
exception was never stamped, so ``async_function_with_retries`` always fell back to the
|
||||
router-wide value. Issue #35071.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def _router(deployment_params: dict, **router_kwargs) -> Router:
|
||||
params = {"model": "openai/gpt-4o", "api_key": "sk-fake"}
|
||||
params.update(deployment_params)
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "mock",
|
||||
"litellm_params": params,
|
||||
"model_info": {"id": "dep-1"},
|
||||
}
|
||||
],
|
||||
**router_kwargs,
|
||||
)
|
||||
|
||||
async def _count_attempts(self, router: Router, **call_kwargs) -> int:
|
||||
calls = {"n": 0}
|
||||
|
||||
async def failing_fn(**kwargs):
|
||||
calls["n"] += 1
|
||||
raise litellm.RateLimitError(message="boom", model="mock", llm_provider="openai")
|
||||
|
||||
with patch("asyncio.sleep", return_value=None):
|
||||
with pytest.raises(litellm.RateLimitError):
|
||||
await router._ageneric_api_call_with_fallbacks(
|
||||
model="mock", original_function=failing_fn, **call_kwargs
|
||||
)
|
||||
return calls["n"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deployment_num_retries_overrides_router_default(self):
|
||||
"""Router-wide 0 + deployment 1 -> 2 attempts; the unfixed path only made 1."""
|
||||
router = self._router({"num_retries": 1}, num_retries=0)
|
||||
assert await self._count_attempts(router, input="hi") == 2
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_num_retries_used_when_deployment_has_none(self):
|
||||
"""Without a deployment override the router-wide value still applies: 1 + 2 = 3."""
|
||||
router = self._router({}, num_retries=2)
|
||||
assert await self._count_attempts(router, input="hi") == 3
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deployment_value_wins_over_request_value_like_chat_completions(self):
|
||||
"""
|
||||
Parity with chat/completions: the stamped deployment value is what the retry loop
|
||||
adopts, even when the request carried its own num_retries (deployment 2 -> 3
|
||||
attempts despite num_retries=0 on the request).
|
||||
"""
|
||||
router = self._router({"num_retries": 2}, num_retries=5)
|
||||
assert await self._count_attempts(router, input="hi", num_retries=0) == 3
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_failed_deployment_id_stamped_on_exception(self):
|
||||
"""
|
||||
The failed deployment id must also be stamped so cooldown/fallback bookkeeping can
|
||||
attribute the failure to the deployment that actually failed.
|
||||
"""
|
||||
router = self._router({"num_retries": 0}, num_retries=0)
|
||||
|
||||
async def failing_fn(**kwargs):
|
||||
raise litellm.RateLimitError(message="boom", model="mock", llm_provider="openai")
|
||||
|
||||
with pytest.raises(litellm.RateLimitError) as exc_info:
|
||||
await router._ageneric_api_call_with_fallbacks(
|
||||
model="mock", original_function=failing_fn, input="hi"
|
||||
)
|
||||
assert getattr(exc_info.value, "failed_deployment_id", None) == "dep-1"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue