Merge pull request #36181 from BerriAI/litellm_fix_batch_group_fallback

fix(router): keep batch fallbacks inside the model group that owns the file
This commit is contained in:
Mateo Wang 2026-08-10 09:39:15 -07:00 • committed by GitHub
commit 60459c60b4
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 242 additions and 3 deletions

View file

@ -1,5 +1,6 @@
import hashlib
import json
from collections.abc import Mapping
from dataclasses import dataclass
from enum import Enum
from typing import TYPE_CHECKING, Any, Final
@ -12,6 +13,7 @@ from litellm.router_utils.add_retry_fallback_headers import (
add_fallback_headers_to_response,
get_fallback_error_info,
)
from litellm.router_utils.batch_utils import _get_router_metadata_variable_name
from litellm.types.router import LiteLLMParamsTypedDict
if TYPE_CHECKING:
@ -131,6 +133,28 @@ def get_fallback_model_group(fallbacks: list[Any], model_group: str) -> tuple[li
return fallback_model_group, generic_fallback_idx
PROVIDER_SCOPED_RESOURCE_KEYS: Final = ("input_file_id", "training_file")
def _get_fallback_target_model_group(fallback_entry: str | Mapping[str, object]) -> str | None:
if isinstance(fallback_entry, str):
return fallback_entry
target: Final = fallback_entry.get("model")
return target if isinstance(target, str) else None
def references_provider_scoped_resource(kwargs: Mapping[str, object]) -> bool:
"""
True when the request names a file that only exists under one provider's credentials.
Batch and fine-tuning jobs are created from a file the caller already uploaded, and
that file lives in the account of the deployment that stored it. Handing the id to a
different model group can only fail, and the second provider's error replaces the
error the caller actually needs to see.
"""
return any(kwargs.get(key) for key in PROVIDER_SCOPED_RESOURCE_KEYS)
async def run_async_fallback(
*args: tuple[Any],
litellm_router: LitellmRouter,
@ -176,6 +200,10 @@ async def run_async_fallback(
error_from_fallbacks = original_exception
fallback_errors = (get_fallback_error_info(original_exception),)
metadata_variable_name: Final = _get_router_metadata_variable_name(
function_name=getattr(kwargs.get("original_function"), "__name__", None)
)
same_model_group_only: Final = references_provider_scoped_resource(kwargs)
# Read out of kwargs and narrowed here rather than declared as a parameter: every caller
# reaches this function by spreading a loosely-typed kwargs dict, so a declared parameter
# would carry an annotation that no call site can actually be checked against.
@ -188,6 +216,13 @@ async def run_async_fallback(
for mg in fallback_model_group:
if mg == original_model_group:
continue
if same_model_group_only and _get_fallback_target_model_group(mg) != original_model_group:
verbose_router_logger.info(
"Skipping fallback to model_group = %s: request is pinned to model_group = %s by its uploaded file",
mask_sensitive_structure(mg),
original_model_group,
)
continue
attempt_key = fallback_attempt_key(mg)
if attempt_key is not None:
if attempt_key in attempted:
@ -205,9 +240,10 @@ async def run_async_fallback(
kwargs["model"] = mg
elif isinstance(mg, dict):
kwargs.update(mg)
kwargs.setdefault("metadata", {}).update(
{"model_group": kwargs.get("model", None)}
) # update model_group used, if fallbacks are done
kwargs[metadata_variable_name] = {
**(kwargs.get(metadata_variable_name) or {}),
"model_group": kwargs.get("model", None),
}
fallback_depth = fallback_depth + 1
kwargs["fallback_depth"] = fallback_depth
kwargs["max_fallbacks"] = max_fallbacks

View file

@ -144,6 +144,147 @@ async def test_run_async_fallback_skips_original_model_group():
assert response._hidden_params["additional_headers"]["x-litellm-attempted-fallbacks"] == 1
class AttemptRecordingRouter:
def __init__(self):
self.attempted_model_groups = []
self.received_kwargs = None
def log_retry(self, kwargs, e):
return kwargs
async def async_function_with_fallbacks(self, *args, **kwargs):
self.attempted_model_groups.append(kwargs.get("model"))
self.received_kwargs = kwargs
return StreamingWrapper()
async def _acreate_batch(*args, **kwargs):
raise AssertionError("only used for its __name__")
@pytest.mark.asyncio
async def test_run_async_fallback_keeps_uploaded_file_requests_in_their_model_group():
"""An input_file_id only exists under the credentials of the group it was uploaded
to, so a cross-group fallback can only fail with the wrong provider's error."""
router = AttemptRecordingRouter()
owning_provider_error = RuntimeError("openai connection error")
with pytest.raises(RuntimeError, match="openai connection error"):
await run_async_fallback(
litellm_router=router,
fallback_model_group=["azure-group"],
original_model_group="openai-group",
original_exception=owning_provider_error,
max_fallbacks=3,
fallback_depth=0,
model="openai-group",
input_file_id="file-owned-by-openai",
original_function=_acreate_batch,
)
assert router.attempted_model_groups == []
@pytest.mark.asyncio
async def test_run_async_fallback_keeps_fine_tuning_requests_in_their_model_group():
router = AttemptRecordingRouter()
with pytest.raises(RuntimeError, match="openai connection error"):
await run_async_fallback(
litellm_router=router,
fallback_model_group=["azure-group"],
original_model_group="openai-group",
original_exception=RuntimeError("openai connection error"),
max_fallbacks=3,
fallback_depth=0,
model="openai-group",
training_file="file-owned-by-openai",
)
assert router.attempted_model_groups == []
@pytest.mark.asyncio
async def test_run_async_fallback_allows_same_model_group_retry_for_uploaded_file_requests():
"""Order-based fallbacks stay inside the owning group, so they must still run."""
router = AttemptRecordingRouter()
await run_async_fallback(
litellm_router=router,
fallback_model_group=[{"model": "openai-group", "_target_order": 2}],
original_model_group="openai-group",
original_exception=RuntimeError("first deployment failed"),
max_fallbacks=3,
fallback_depth=0,
model="openai-group",
input_file_id="file-owned-by-openai",
original_function=_acreate_batch,
)
assert router.attempted_model_groups == ["openai-group"]
@pytest.mark.asyncio
async def test_run_async_fallback_still_crosses_model_groups_without_an_uploaded_file():
router = AttemptRecordingRouter()
await run_async_fallback(
litellm_router=router,
fallback_model_group=["azure-group"],
original_model_group="openai-group",
original_exception=RuntimeError("openai connection error"),
max_fallbacks=3,
fallback_depth=0,
model="openai-group",
)
assert router.attempted_model_groups == ["azure-group"]
@pytest.mark.asyncio
async def test_run_async_fallback_handles_explicitly_none_metadata():
"""/v1/batches always sets `metadata`, and sets it to None when the caller sent
none, so setdefault() on it hands back None instead of a dict."""
router = AttemptRecordingRouter()
await run_async_fallback(
litellm_router=router,
fallback_model_group=["azure-group"],
original_model_group="openai-group",
original_exception=RuntimeError("openai connection error"),
max_fallbacks=3,
fallback_depth=0,
model="openai-group",
metadata=None,
)
assert router.received_kwargs["metadata"] == {"model_group": "azure-group"}
@pytest.mark.asyncio
async def test_run_async_fallback_records_batch_model_group_outside_provider_metadata():
"""`metadata` on a batch request is forwarded to the provider and stored on the
batch, so the router's own model_group belongs in litellm_metadata."""
router = AttemptRecordingRouter()
await run_async_fallback(
litellm_router=router,
fallback_model_group=[{"model": "openai-group", "_target_order": 2}],
original_model_group="openai-group",
original_exception=RuntimeError("first deployment failed"),
max_fallbacks=3,
fallback_depth=0,
model="openai-group",
input_file_id="file-owned-by-openai",
metadata={"caller": "nightly-job"},
litellm_metadata={"model_group": "openai-group"},
original_function=_acreate_batch,
)
assert router.received_kwargs["metadata"] == {"caller": "nightly-job"}
assert router.received_kwargs["litellm_metadata"]["model_group"] == "openai-group"
class RecordingFailRouter:
def __init__(self):
self.attempted_models = []

View file

@ -6757,6 +6757,68 @@ async def test_acreate_batch_disable_fallbacks_surfaces_owning_provider_error():
assert mock_create.call_args.kwargs["model"] == "owning-model"
@pytest.mark.asyncio
async def test_acreate_batch_surfaces_owning_provider_error_without_disable_fallbacks():
"""The router itself has to keep a batch inside the group that owns the input file:
the proxy only sets disable_fallbacks on the managed-files route, so the caller
otherwise gets the fallback provider's error for a file it never received."""
from litellm.types.utils import LiteLLMBatch
router = litellm.Router(
model_list=[
{
"model_name": "owning-model",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "sk-owning",
},
},
{
"model_name": "fallback-model",
"litellm_params": {
"model": "azure/gpt-4o-mini",
"api_key": "sk-fallback",
"api_base": "https://fallback.openai.azure.com",
"api_version": "2024-08-01-preview",
},
},
],
fallbacks=[{"owning-model": ["fallback-model"]}],
num_retries=0,
)
attempted_models = []
async def _acreate_batch(model, **kwargs):
attempted_models.append(model)
if model == "owning-model":
raise litellm.APIConnectionError(
message="Connection error - openai is unreachable",
model="openai/gpt-4o-mini",
llm_provider="openai",
)
return LiteLLMBatch(
id="batch-created-on-the-wrong-provider",
completion_window="24h",
created_at=0,
endpoint="/v1/chat/completions",
input_file_id="file-owned-by-openai",
object="batch",
status="validating",
)
with patch.object(router, "_acreate_batch", _acreate_batch):
with pytest.raises(litellm.APIConnectionError, match="openai is unreachable"):
await router.acreate_batch(
model="owning-model",
input_file_id="file-owned-by-openai",
endpoint="/v1/chat/completions",
completion_window="24h",
metadata={"team": "batch-jobs"},
)
assert attempted_models == ["owning-model"]
@pytest.mark.asyncio
async def test_acreate_batch_request_bedrock_tags_override_deployment_tags():
import httpx