fix(router): handle num_retries=None to prevent TypeError in retry logic

When num_retries is passed as None (either explicitly by the caller or
through certain proxy configurations), the comparison `if num_retries > 0`
in async_function_with_retries raises:
  TypeError: '>' not supported between instances of 'NoneType' and 'int'

The root cause is that `kwargs.get("num_retries", self.num_retries)` returns
None when the key exists in kwargs with an explicit None value, bypassing
the default fallback to self.num_retries.

This fix ensures num_retries is never None by:
1. Using an explicit None check when setting kwargs["num_retries"] across
   all 8 call sites (acompletion, image_generation, text_completion, etc.)
2. Adding a None guard in async_function_with_retries where num_retries
   is popped from kwargs, falling back to self.num_retries

Fixes #23316

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
gambletan 2026-03-11 10:25:14 +08:00
parent 82a9b0ea03
commit b5a733bf8a

View file

@ -2108,7 +2108,7 @@ class Router:
- litellm_trace_id
- metadata
"""
kwargs["num_retries"] = kwargs.get("num_retries", self.num_retries)
kwargs["num_retries"] = kwargs.get("num_retries") if kwargs.get("num_retries") is not None else self.num_retries
kwargs.setdefault("litellm_trace_id", str(uuid.uuid4()))
model_group_alias: Optional[str] = None
if self._get_model_from_alias(model=model):
@ -2848,7 +2848,7 @@ class Router:
kwargs["model"] = model
kwargs["prompt"] = prompt
kwargs["original_function"] = self._image_generation
kwargs["num_retries"] = kwargs.get("num_retries", self.num_retries)
kwargs["num_retries"] = kwargs.get("num_retries") if kwargs.get("num_retries") is not None else self.num_retries
kwargs.setdefault("metadata", {}).update({"model_group": model})
response = self.function_with_fallbacks(**kwargs)
@ -2907,7 +2907,7 @@ class Router:
kwargs["model"] = model
kwargs["prompt"] = prompt
kwargs["original_function"] = self._aimage_generation
kwargs["num_retries"] = kwargs.get("num_retries", self.num_retries)
kwargs["num_retries"] = kwargs.get("num_retries") if kwargs.get("num_retries") is not None else self.num_retries
self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs)
response = await self.async_function_with_fallbacks(**kwargs)
@ -3270,7 +3270,7 @@ class Router:
try:
kwargs["model"] = model
kwargs["prompt"] = prompt
kwargs["num_retries"] = kwargs.get("num_retries", self.num_retries)
kwargs["num_retries"] = kwargs.get("num_retries") if kwargs.get("num_retries") is not None else self.num_retries
kwargs.setdefault("metadata", {}).update({"model_group": model})
# pick the one that is available (lowest TPM/RPM)
@ -3414,7 +3414,7 @@ class Router:
kwargs["model"] = model
kwargs["adapter_id"] = adapter_id
kwargs["original_function"] = self._aadapter_completion
kwargs["num_retries"] = kwargs.get("num_retries", self.num_retries)
kwargs["num_retries"] = kwargs.get("num_retries") if kwargs.get("num_retries") is not None else self.num_retries
kwargs.setdefault("metadata", {}).update({"model_group": model})
response = await self.async_function_with_fallbacks(**kwargs)
@ -4011,7 +4011,7 @@ class Router:
try:
kwargs["model"] = model
kwargs["original_function"] = self._acreate_file
kwargs["num_retries"] = kwargs.get("num_retries", self.num_retries)
kwargs["num_retries"] = kwargs.get("num_retries") if kwargs.get("num_retries") is not None else self.num_retries
self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs)
response = await self.async_function_with_fallbacks(**kwargs)
@ -4280,7 +4280,7 @@ class Router:
try:
kwargs["model"] = model
kwargs["original_function"] = self._acreate_batch
kwargs["num_retries"] = kwargs.get("num_retries", self.num_retries)
kwargs["num_retries"] = kwargs.get("num_retries") if kwargs.get("num_retries") is not None else self.num_retries
metadata_variable_name = _get_router_metadata_variable_name(
function_name="_acreate_batch"
)
@ -4514,7 +4514,7 @@ class Router:
try:
kwargs["model"] = model
kwargs["original_function"] = self._acancel_batch
kwargs["num_retries"] = kwargs.get("num_retries", self.num_retries)
kwargs["num_retries"] = kwargs.get("num_retries") if kwargs.get("num_retries") is not None else self.num_retries
metadata_variable_name = _get_router_metadata_variable_name(
function_name="_acancel_batch"
)
@ -5383,7 +5383,9 @@ class Router:
"model_group_retry_policy", self.model_group_retry_policy
)
model_group: Optional[str] = kwargs.get("model")
num_retries = kwargs.pop("num_retries")
num_retries = kwargs.pop("num_retries", self.num_retries)
if num_retries is None:
num_retries = self.num_retries
## ADD MODEL GROUP SIZE TO METADATA - used for model_group_rate_limit_error tracking
_metadata: dict = kwargs.get("litellm_metadata", kwargs.get("metadata")) or {}