mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
commit
60459c60b4
3 changed files with 242 additions and 3 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 = []
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue