fix(proxy): preserve blocked-model fallback validation across retries

This commit is contained in:
Duxl-Ai 2026-08-27 04:25:38 +08:00
parent f6571a653f
commit 8ab1860051
8 changed files with 1593 additions and 186 deletions

View file

@ -7,7 +7,15 @@ from fastapi import HTTPException, status
import litellm
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
from litellm.router_utils.common_utils import _is_proxy_admin_request
from litellm.router_utils.common_utils import _is_proxy_admin_request, resolve_model_group_alias
from litellm.router_utils.fallback_event_handlers import (
_check_non_standard_fallback_format,
creates_provider_scoped_resource,
get_authenticated_team_context,
get_fallback_model_group,
preserve_authenticated_team_context,
references_provider_scoped_resource,
)
# Client-supplied params that make the router or the call path fabricate a
# failure or a delay instead of calling the provider. The ``mock_testing_*``
@ -27,6 +35,22 @@ GATED_MOCK_PARAM_NAMES: Final[tuple[str, ...]] = (
MOCK_TESTING_CONFIG_KEY: Final = "dangerously_allow_mock_testing_request_params"
EVAL_ROUTE_TYPES: Final = frozenset(
{
"acreate_eval",
"alist_evals",
"aget_eval",
"aupdate_eval",
"adelete_eval",
"acancel_eval",
"acreate_run",
"alist_runs",
"aget_run",
"acancel_run",
"adelete_run",
}
)
if TYPE_CHECKING:
from litellm.router import Router as _Router
@ -55,13 +79,244 @@ def _is_a2a_agent_model(model_name: Any) -> bool:
return isinstance(model_name, str) and model_name.startswith("a2a/")
def _raise_if_model_fully_blocked(llm_router: LitellmRouter, model_name: Any, team_id: str | None) -> None:
def _deployment_ids(deployments: Any) -> set[str]:
"""Return deployment IDs from either Router healthy-deployment shape."""
deployment_list = [deployments] if isinstance(deployments, Mapping) else deployments
if not isinstance(deployment_list, list):
return set()
deployment_ids: set[str] = set()
for deployment in deployment_list:
if not isinstance(deployment, Mapping):
continue
model_info = deployment.get("model_info")
if not isinstance(model_info, Mapping):
continue
deployment_id = model_info.get("id")
if isinstance(deployment_id, str) and deployment_id:
deployment_ids.add(deployment_id)
return deployment_ids
def _expand_fallback_target(fallback_target: Any) -> list[str | dict[str, Any]]:
"""Expand a list-valued direct fallback into ordered concrete candidates."""
if isinstance(fallback_target, str):
return [fallback_target]
if not isinstance(fallback_target, Mapping):
return []
fallback_model = fallback_target.get("model")
if isinstance(fallback_model, str):
return [dict(fallback_target)]
if not isinstance(fallback_model, list) or not fallback_model:
return []
if not all(isinstance(candidate, str) and candidate for candidate in fallback_model):
return []
return [{**fallback_target, "model": candidate} for candidate in fallback_model]
def _expand_fallback_tail(fallback_targets: list[Any]) -> list[Any]:
"""Expand list-valued direct targets before handing the trusted tail to runtime."""
expanded_targets: list[Any] = []
for fallback_target in fallback_targets:
candidates = _expand_fallback_target(fallback_target)
if candidates:
expanded_targets.extend(candidates)
else:
expanded_targets.append(fallback_target)
return expanded_targets
async def _get_available_fallback_request(
llm_router: LitellmRouter,
model_name: str,
team_id: str | None,
request_data: Mapping[str, Any],
) -> dict[str, Any] | None:
"""Return a trusted request targeting the first executable server fallback."""
if request_data.get("disable_fallbacks") is True:
return None
max_fallbacks: Final = request_data.get("max_fallbacks", getattr(llm_router, "max_fallbacks", None))
if max_fallbacks == 0:
return None
fallbacks: Final = getattr(llm_router, "fallbacks", None)
if not isinstance(fallbacks, list):
return None
router_model_list: Final = getattr(llm_router, "model_list", None)
if not isinstance(router_model_list, list):
return None
router_deployment_ids: Final = _deployment_ids(router_model_list)
# Runtime enforcement below is fail-closed: every Router deployment must have
# an ID so an unvalidated target cannot escape the exclusion boundary.
if len(router_deployment_ids) != len(router_model_list):
return None
if _check_non_standard_fallback_format(fallbacks):
fallback_model_group = fallbacks
else:
fallback_model_group, _ = get_fallback_model_group(
fallbacks=fallbacks,
model_group=model_name,
)
if not isinstance(fallback_model_group, list):
return None
same_model_group_only: Final = references_provider_scoped_resource(
request_data
) or creates_provider_scoped_resource(request_data)
alias_map: Final = getattr(llm_router, "model_group_alias", None)
canonical_model_group: Final = resolve_model_group_alias(alias_map, model_name) or model_name
_, authenticated_team_bucket = get_authenticated_team_context(request_data)
for fallback_index, fallback_target in enumerate(fallback_model_group):
fallback_candidates = _expand_fallback_target(fallback_target)
for candidate_index, candidate_target in enumerate(fallback_candidates):
# Both copies start without caller-supplied fallbacks. ``fallback_request``
# is disposable preflight state; ``runtime_request`` is the trusted request
# we will hand to the Router if this server-configured target is eligible.
runtime_request: dict[str, Any] = dict(request_data)
runtime_request.pop("fallbacks", None)
fallback_request: dict[str, Any] = dict(runtime_request)
if isinstance(candidate_target, str):
fallback_model = candidate_target
canonical_fallback_model = resolve_model_group_alias(alias_map, fallback_model) or fallback_model
if canonical_fallback_model == canonical_model_group:
continue
else:
fallback_model = candidate_target["model"]
canonical_fallback_model = resolve_model_group_alias(alias_map, fallback_model) or fallback_model
fallback_request.update(candidate_target)
runtime_request.update(candidate_target)
if same_model_group_only and canonical_fallback_model != canonical_model_group:
continue
fallback_request["model"] = fallback_model
runtime_request["model"] = fallback_model
preserve_authenticated_team_context(
request_kwargs=fallback_request,
authenticated_team_id=team_id,
source_bucket=authenticated_team_bucket,
)
preserve_authenticated_team_context(
request_kwargs=runtime_request,
authenticated_team_id=team_id,
source_bucket=authenticated_team_bucket,
)
try:
fallback_messages = fallback_request.get("messages")
fallback_input = fallback_request.get("input")
pre_routing_hook_response = await llm_router.async_pre_routing_hook(
model=fallback_model,
request_kwargs=fallback_request,
messages=fallback_messages,
input=fallback_input,
specific_deployment=False,
)
if pre_routing_hook_response is not None:
fallback_model = pre_routing_hook_response.model
fallback_messages = pre_routing_hook_response.messages
if pre_routing_hook_response.litellm_params is not None:
fallback_request.update(pre_routing_hook_response.litellm_params)
preserve_authenticated_team_context(
request_kwargs=fallback_request,
authenticated_team_id=team_id,
source_bucket=authenticated_team_bucket,
)
healthy_deployments = await llm_router.async_get_healthy_deployments(
model=fallback_model,
request_kwargs=fallback_request,
messages=fallback_messages,
input=fallback_input,
)
except Exception:
continue
validated_deployment_ids: Final = _deployment_ids(healthy_deployments)
if not validated_deployment_ids:
continue
# ``async_get_available_deployment`` invokes the pre-routing hook again at
# runtime. A stateful strategy can therefore return a different target on
# the second invocation. Router already supports a one-shot exclusion list
# that is consumed by the next healthy-deployment lookup; restrict that
# lookup to deployments that actually passed this preflight validation.
existing_exclusions = runtime_request.get("_excluded_deployment_ids")
excluded_deployment_ids = (
{deployment_id for deployment_id in existing_exclusions if isinstance(deployment_id, str)}
if isinstance(existing_exclusions, (list, tuple, set, frozenset))
else set()
)
excluded_deployment_ids.update(router_deployment_ids - validated_deployment_ids)
runtime_request["_excluded_deployment_ids"] = sorted(excluded_deployment_ids)
# We are consuming the first fallback hop here instead of letting the
# blocked primary enter normal Router selection. Preserve the remaining
# candidates from a list-valued target before later trusted fallbacks,
# unless the selected server-side dict supplied its own fallback chain.
if "fallbacks" not in runtime_request:
remaining_candidates = fallback_candidates[candidate_index + 1 :]
later_fallbacks = _expand_fallback_tail(fallback_model_group[fallback_index + 1 :])
runtime_request["fallbacks"] = [*remaining_candidates, *later_fallbacks]
fallback_depth = request_data.get("fallback_depth", 0)
runtime_request["fallback_depth"] = fallback_depth + 1 if isinstance(fallback_depth, int) else 1
return runtime_request
return None
async def _has_available_fallback(
llm_router: LitellmRouter,
model_name: str,
team_id: str | None,
request_data: Mapping[str, Any],
) -> bool:
"""Return whether a server-configured fallback is eligible for this request at runtime."""
return (
await _get_available_fallback_request(
llm_router=llm_router,
model_name=model_name,
team_id=team_id,
request_data=request_data,
)
is not None
)
async def _raise_if_model_fully_blocked(
llm_router: LitellmRouter,
model_name: Any,
team_id: str | None,
request_data: dict[str, Any],
*,
allow_router_fallback: bool = True,
) -> None:
if not isinstance(model_name, str) or not model_name:
return
if not isinstance(llm_router, litellm.Router):
return
deployments: Final = llm_router.get_model_list(model_name=model_name, team_id=team_id) or []
if llm_router._are_all_deployments_blocked(deployments):
if not llm_router._are_all_deployments_blocked(deployments):
return
fallback_request = (
await _get_available_fallback_request(
llm_router=llm_router,
model_name=model_name,
team_id=team_id,
request_data=request_data,
)
if allow_router_fallback
else None
)
if fallback_request is None:
raise litellm.PermissionDeniedError(
message="Model is blocked",
model=model_name,
@ -72,6 +327,13 @@ def _raise_if_model_fully_blocked(llm_router: LitellmRouter, model_name: Any, te
),
)
# Never send the fully blocked primary back through normal Router selection:
# a request-dependent pre-routing strategy/plugin could otherwise rewrite it
# into a non-fallback model that was never authorized. The request now targets
# the first executable server-configured fallback validated above.
request_data.clear()
request_data.update(fallback_request)
ROUTE_ENDPOINT_MAPPING: Final = {
"acompletion": "/chat/completions",
@ -528,29 +790,22 @@ async def _route_request_single_attempt( # noqa: ANN202 # returns unawaited pr
for key in per_request_settings:
if key in override_settings and key not in data:
data[key] = override_settings[key]
# Use main router with overridden kwargs
if llm_router is not None:
return getattr(llm_router, f"{route_type}")(**data)
else:
return getattr(litellm, f"{route_type}")(**data)
elif llm_router is not None:
_raise_if_model_fully_blocked(llm_router=llm_router, model_name=data.get("model"), team_id=team_id)
await _raise_if_model_fully_blocked(
llm_router=llm_router,
model_name=data.get("model"),
team_id=team_id,
request_data=data,
allow_router_fallback=route_type not in EVAL_ROUTE_TYPES,
)
# Evals API: always route to litellm directly (not through router)
# But extract model credentials if a model is provided
if route_type in [
"acreate_eval",
"alist_evals",
"aget_eval",
"aupdate_eval",
"adelete_eval",
"acancel_eval",
"acreate_run",
"alist_runs",
"aget_run",
"acancel_run",
"adelete_run",
]:
if route_type in EVAL_ROUTE_TYPES:
# If a model is provided, get its credentials from the router
model: Final = data.get("model")
if model and llm_router:

View file

@ -14,6 +14,8 @@ from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.types.router import CredentialLiteLLMParams
from litellm.types.utils import LlmProviders
_RETRY_SCOPED_EXCLUSION_STATE_KEY: Final = "_retry_scoped_excluded_deployment_ids"
def _is_proxy_admin_request(request_kwargs: Mapping[str, object] | None) -> bool:
if request_kwargs is None:
@ -103,7 +105,11 @@ def filter_team_based_models(
"""
If a model has a team_id
Only use if request is from that team
Only use if request is from that team. Router-internal deployment exclusions
are also enforced here so the single-dict specific-deployment shape cannot
bypass the same exclusion boundary applied to model-group lists. Exclusions
persist across retries of the same target and are cleared only when fallback
execution advances to the next trusted target.
"""
if request_kwargs is None:
return healthy_deployments
@ -111,7 +117,56 @@ def filter_team_based_models(
metadata: Final = request_kwargs.get("metadata") or {}
litellm_metadata: Final = request_kwargs.get("litellm_metadata") or {}
request_team_id: Final = metadata.get("user_api_key_team_id") or litellm_metadata.get("user_api_key_team_id")
if request_team_id is None and _is_proxy_admin_request(request_kwargs) and isinstance(healthy_deployments, list):
# Router health selection consumes ``_excluded_deployment_ids`` after this
# filter runs. Keep a retry-scoped copy keyed by fallback depth so every
# retry of the same trusted target restores the preflight boundary, while
# advancing to the next fallback depth naturally discards the old state.
fallback_depth: Final = request_kwargs.get("fallback_depth")
raw_excluded_deployment_ids = request_kwargs.get("_excluded_deployment_ids")
retry_state: Final = request_kwargs.get(_RETRY_SCOPED_EXCLUSION_STATE_KEY)
if isinstance(raw_excluded_deployment_ids, (list, tuple, set, frozenset)):
persisted_ids: Final = [
deployment_id for deployment_id in raw_excluded_deployment_ids if isinstance(deployment_id, str)
]
request_kwargs[_RETRY_SCOPED_EXCLUSION_STATE_KEY] = {
"fallback_depth": fallback_depth,
"deployment_ids": persisted_ids,
}
elif isinstance(retry_state, Mapping) and retry_state.get("fallback_depth") == fallback_depth:
persisted_ids_value: Final = retry_state.get("deployment_ids")
if isinstance(persisted_ids_value, (list, tuple, set, frozenset)):
raw_excluded_deployment_ids = [
deployment_id for deployment_id in persisted_ids_value if isinstance(deployment_id, str)
]
request_kwargs["_excluded_deployment_ids"] = list(raw_excluded_deployment_ids)
elif retry_state is not None:
request_kwargs.pop(_RETRY_SCOPED_EXCLUSION_STATE_KEY, None)
excluded_deployment_ids: Final = (
{deployment_id for deployment_id in raw_excluded_deployment_ids if isinstance(deployment_id, str)}
if isinstance(raw_excluded_deployment_ids, (list, tuple, set, frozenset))
else set()
)
# A specific deployment ID is returned as a single dict instead of a list.
# Apply both exclusion and team isolation here rather than treating the shape
# as an implicit authorization bypass. Proxy admins retain their existing
# ability to address a team-scoped deployment directly, but cannot override
# an explicit Router-internal exclusion boundary.
if isinstance(healthy_deployments, dict):
model_info: Final = healthy_deployments.get("model_info") or {}
deployment_id: Final = model_info.get("id")
if deployment_id in excluded_deployment_ids:
return []
model_team_id: Final = model_info.get("team_id")
if model_team_id is None or model_team_id == request_team_id:
return healthy_deployments
if request_team_id is None and _is_proxy_admin_request(request_kwargs):
return healthy_deployments
return []
if request_team_id is None and _is_proxy_admin_request(request_kwargs):
requested_model: Final = (
request_kwargs.get("model") or metadata.get("model_group") or litellm_metadata.get("model_group")
)
@ -143,11 +198,13 @@ def filter_team_based_models(
llm_provider="",
)
if matches_requested_model:
return healthy_deployments
return [
deployment
for deployment in healthy_deployments
if deployment.get("model_info", {}).get("id") not in excluded_deployment_ids
]
ids_to_remove: Final = set()
if isinstance(healthy_deployments, dict):
return healthy_deployments
ids_to_remove: Final = set(excluded_deployment_ids)
for deployment in healthy_deployments:
_model_info = deployment.get("model_info") or {}
model_team_id = _model_info.get("team_id")
@ -165,7 +222,7 @@ def filter_team_based_models(
def _deployment_supports_web_search(deployment: dict) -> bool:
"""
Check if a deployment supports web search.
Check if a deployment supports web search
Priority:
1. Check config-level override in model_info.supports_web_search
@ -208,18 +265,12 @@ def filter_web_search_deployments(
if not is_web_search_request:
return healthy_deployments
# Filter out deployments that don't support web search
final_deployments: Final = [d for d in healthy_deployments if _deployment_supports_web_search(d)]
if len(healthy_deployments) > 0 and len(final_deployments) == 0:
verbose_logger.warning("No deployments support web search for request")
return final_deployments
# Credential params that only one provider family reads, paired with the providers
# that read them. A deployment carrying them while resolving elsewhere is almost
# always a missing route prefix: `model: claude-sonnet-5` with `aws_region_name`
# set resolves to the first-party Anthropic API, silently ignores the AWS
# credentials, and 401s at request time.
_AWS_PROVIDERS: Final = frozenset(
provider.value for provider in LlmProviders if provider.value.startswith(("bedrock", "sagemaker"))
)

View file

@ -14,6 +14,7 @@ from litellm.router_utils.add_retry_fallback_headers import (
get_fallback_error_info,
)
from litellm.router_utils.batch_utils import _get_router_metadata_variable_name
from litellm.router_utils.common_utils import resolve_model_group_alias
from litellm.router_utils.cooldown_handlers import (
_first_present, # pyright: ignore[reportPrivateUsage] - shared internal helper, used across router_utils
_set_cooldown_deployments, # pyright: ignore[reportPrivateUsage] - shared helper, used across router_utils
@ -35,6 +36,101 @@ else:
# Status codes a generic API call's caller-supplied resource id can trigger on its own
# (e.g. a nonexistent file/batch/thread id), independent of the selected deployment's health.
_REQUEST_SCOPED_STATUS_CODES: Final = frozenset((404,))
_ROUTER_METADATA_BUCKETS: Final = ("metadata", "litellm_metadata")
_TEAM_ID_METADATA_KEY: Final = "user_api_key_team_id"
_API_KEY_AUTH_METADATA_KEY: Final = "user_api_key_auth"
@dataclass(frozen=True, slots=True)
class AuthenticatedMetadataContext:
"""Proxy-authenticated metadata that fallback overrides must not replace."""
team_source_bucket: str | None = None
api_key_auth_source_bucket: str | None = None
user_api_key_auth: Any = None
has_user_api_key_auth: bool = False
def get_authenticated_team_context(
request_kwargs: Mapping[str, Any],
) -> tuple[str | None, AuthenticatedMetadataContext]:
"""Return authenticated team and API-key context before fallback overrides."""
authenticated_team_id: str | None = None
team_source_bucket: str | None = None
api_key_auth_source_bucket: str | None = None
user_api_key_auth: Any = None
has_user_api_key_auth = False
for bucket_name in _ROUTER_METADATA_BUCKETS:
bucket = request_kwargs.get(bucket_name)
if not isinstance(bucket, Mapping):
continue
if authenticated_team_id is None:
team_id = bucket.get(_TEAM_ID_METADATA_KEY)
if isinstance(team_id, str):
authenticated_team_id = team_id
team_source_bucket = bucket_name
if not has_user_api_key_auth and _API_KEY_AUTH_METADATA_KEY in bucket:
user_api_key_auth = bucket.get(_API_KEY_AUTH_METADATA_KEY)
api_key_auth_source_bucket = bucket_name
has_user_api_key_auth = True
return authenticated_team_id, AuthenticatedMetadataContext(
team_source_bucket=team_source_bucket,
api_key_auth_source_bucket=api_key_auth_source_bucket,
user_api_key_auth=user_api_key_auth,
has_user_api_key_auth=has_user_api_key_auth,
)
def preserve_authenticated_team_context(
request_kwargs: dict[str, Any],
authenticated_team_id: str | None,
source_bucket: AuthenticatedMetadataContext | str | None,
) -> None:
"""Keep proxy-authenticated team/API-key metadata authoritative across fallbacks."""
context = (
source_bucket
if isinstance(source_bucket, AuthenticatedMetadataContext)
else AuthenticatedMetadataContext(team_source_bucket=source_bucket)
)
for bucket_name in _ROUTER_METADATA_BUCKETS:
bucket = request_kwargs.get(bucket_name)
if not isinstance(bucket, Mapping):
continue
updated_bucket = dict(bucket)
if authenticated_team_id is None:
updated_bucket.pop(_TEAM_ID_METADATA_KEY, None)
else:
updated_bucket[_TEAM_ID_METADATA_KEY] = authenticated_team_id
if context.has_user_api_key_auth:
updated_bucket[_API_KEY_AUTH_METADATA_KEY] = context.user_api_key_auth
else:
updated_bucket.pop(_API_KEY_AUTH_METADATA_KEY, None)
request_kwargs[bucket_name] = updated_bucket
if authenticated_team_id is not None:
authoritative_bucket = (
context.team_source_bucket
if context.team_source_bucket in _ROUTER_METADATA_BUCKETS
else "metadata"
)
bucket = request_kwargs.get(authoritative_bucket)
updated_bucket = dict(bucket) if isinstance(bucket, Mapping) else {}
updated_bucket[_TEAM_ID_METADATA_KEY] = authenticated_team_id
request_kwargs[authoritative_bucket] = updated_bucket
if context.has_user_api_key_auth:
authoritative_bucket = (
context.api_key_auth_source_bucket
if context.api_key_auth_source_bucket in _ROUTER_METADATA_BUCKETS
else "metadata"
)
bucket = request_kwargs.get(authoritative_bucket)
updated_bucket = dict(bucket) if isinstance(bucket, Mapping) else {}
updated_bucket[_API_KEY_AUTH_METADATA_KEY] = context.user_api_key_auth
request_kwargs[authoritative_bucket] = updated_bucket
def _trigger_cooldown_for_failed_deployment(
@ -60,10 +156,6 @@ def _trigger_cooldown_for_failed_deployment(
exception_status: Final[str | int] = getattr(exception, "status_code", "")
# Generic API calls (files, batches, threads, rerank, ...) take a caller-supplied
# resource id, so a 404 there usually means "that id doesn't exist" rather than
# "this deployment is unhealthy". Left unguarded, one bad id would 404 every
# deployment in the fallback chain and cool all of them down from a single request.
if (
kwargs.get("original_generic_function") is not None
and cast_exception_status_to_int(exception_status) in _REQUEST_SCOPED_STATUS_CODES
@ -75,10 +167,6 @@ def _trigger_cooldown_for_failed_deployment(
)
return
# The proxy's `x-litellm-timeout` header lets a caller set an arbitrarily short
# timeout, which litellm.Timeout reports as status 408 regardless of the deployment's
# actual health. Left unguarded, a caller could force a 408 on every deployment in
# the fallback chain from a single request with a near-zero timeout.
if kwargs.get("client_side_timeout") and cast_exception_status_to_int(exception_status) == 408:
verbose_router_logger.debug(
"Not triggering cooldown for fallback deployment: a caller-supplied "
@ -86,19 +174,12 @@ def _trigger_cooldown_for_failed_deployment(
)
return
# Only Router._set_failed_deployment_id_on_exception()'s server-stamped id is
# trusted here: a metadata-bucket lookup (e.g. "metadata"/"litellm_metadata")
# can't reliably tell a caller-supplied bucket from a router-authored one
# without knowing this call's function_name, so a client with permission to
# set metadata could otherwise get an arbitrary deployment cooled down.
deployment_id: Final[str | None] = getattr(exception, "failed_deployment_id", None)
if deployment_id is None:
verbose_router_logger.debug("Cannot trigger cooldown for fallback: no failed_deployment_id on exception")
return
# Priority: deployment config > response header > router default, matching
# Router.deployment_callback_on_failure's precedence for the primary path.
deployment_dict: Final = litellm_router.get_model_info(id=deployment_id)
deployment_cooldown: Final = (
_first_present(
@ -170,15 +251,7 @@ def fallback_attempt_key(fallback_target: object) -> str | None:
@dataclass(slots=True)
class AttemptedFallbackTargets:
"""
The fallback attempts a single request has already made.
One instance is created on the first fallback hop and shared by reference for the rest
of the walk, so an attempt made in one branch is not repeated in a sibling branch.
Without it the walk enumerates paths rather than attempts: a fallback graph containing
a cycle retries one deterministic failure once per path through the cycle, and a
client-side fallback list is re-walked at every level of the recursion.
"""
"""The fallback attempts a single request has already made."""
keys: frozenset[str] = frozenset()
@ -190,18 +263,6 @@ class AttemptedFallbackTargets:
def _check_stripped_model_group(model_group: str, fallback_key: str) -> bool:
"""
Handles wildcard routing scenario
where fallbacks set like:
[{"gpt-3.5-turbo": ["claude-3-haiku"]}]
but model_group is like:
"openai/gpt-3.5-turbo"
Returns:
- True if the stripped model group == fallback_key
"""
for provider in litellm.provider_list:
if isinstance(provider, Enum):
_provider = provider.value
@ -215,34 +276,20 @@ def _check_stripped_model_group(model_group: str, fallback_key: str) -> bool:
def get_fallback_model_group(fallbacks: list[Any], model_group: str) -> tuple[list[str] | None, int | None]:
"""
Returns:
- fallback_model_group: List[str] of fallback model groups. example: ["gpt-4", "gpt-3.5-turbo"]
- generic_fallback_idx: int of the index of the generic fallback in the fallbacks list.
Checks:
- exact match
- stripped model group match
- generic fallback
"""
generic_fallback_idx: int | None = None
stripped_model_fallback: list[str] | None = None
fallback_model_group: list[str] | None = None
## check for specific model group-specific fallbacks
for idx, item in enumerate(fallbacks):
if isinstance(item, dict):
if list(item.keys())[0] == model_group: # check exact match
if list(item.keys())[0] == model_group:
fallback_model_group = item[model_group]
break
elif _check_stripped_model_group(
model_group=model_group, fallback_key=list(item.keys())[0]
): # check generic fallback
elif _check_stripped_model_group(model_group=model_group, fallback_key=list(item.keys())[0]):
stripped_model_fallback = item[list(item.keys())[0]]
elif list(item.keys())[0] == "*": # check generic fallback
elif list(item.keys())[0] == "*":
generic_fallback_idx = idx
elif isinstance(item, str):
fallback_model_group = [item]
## if none, check for generic fallback
if fallback_model_group is None:
if stripped_model_fallback is not None:
fallback_model_group = stripped_model_fallback
@ -264,26 +311,10 @@ def _get_fallback_target_model_group(fallback_entry: str | Mapping[str, object])
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)
def creates_provider_scoped_resource(kwargs: Mapping[str, object]) -> bool:
"""
True when the request creates a resource that will live under one provider's credentials.
A file uploaded for batches or fine-tuning is stored in the account of the deployment
that handled it, and its id is only usable against the model group the caller named.
Letting the upload fall back to a different model group silently stores the file with
the wrong provider, and every later use of the returned id fails.
"""
return getattr(kwargs.get("original_function"), "__name__", None) in PROVIDER_SCOPED_CREATION_FUNCTION_NAMES
@ -298,35 +329,6 @@ async def run_async_fallback(
include_fallback_errors: bool = False,
**kwargs,
) -> Any:
"""
Loops through all the fallback model groups and calls kwargs["original_function"] with the arguments and keyword arguments provided.
If the call is successful, it logs the success and returns the response.
If the call fails, it logs the failure and continues to the next fallback model group.
If all fallback model groups fail, it raises the most recent exception.
Args:
litellm_router: The litellm router instance.
*args: Positional arguments.
fallback_model_group: List[str] of fallback model groups. example: ["gpt-4", "gpt-3.5-turbo"]
original_model_group: The original model group. example: "gpt-3.5-turbo"
original_exception: The original exception.
**kwargs: Keyword arguments. `attempted_targets` carries the fallback attempts
already made for this request, created on the first hop and shared by reference
for the rest of the walk. A target already in it is skipped, so neither a
fallback graph that loops back on itself nor a client-side fallback list
re-walked at each level can repeat an attempt that has already failed. Identity
comes from `fallback_attempt_key`, so an entry that overrides request params or
re-targets the failed group with a different deployment selection stays distinct
from a bare name.
Returns:
The response from the successful fallback model group.
Raises:
The most recent exception if all fallback model groups fail.
"""
### BASE CASE ### MAX FALLBACK DEPTH REACHED
if fallback_depth >= max_fallbacks:
raise original_exception
@ -335,12 +337,14 @@ async def run_async_fallback(
metadata_variable_name: Final = _get_router_metadata_variable_name(
function_name=getattr(kwargs.get("original_function"), "__name__", None)
)
authenticated_team_id, authenticated_team_bucket = get_authenticated_team_context(kwargs)
same_model_group_only: Final = references_provider_scoped_resource(kwargs) or creates_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.
alias_map: Final = getattr(litellm_router, "model_group_alias", None)
canonical_original_model_group: Final = (
resolve_model_group_alias(alias_map, original_model_group) or original_model_group
)
carried_targets: Final = kwargs.get("attempted_targets")
attempted: Final = (
carried_targets if isinstance(carried_targets, AttemptedFallbackTargets) else AttemptedFallbackTargets()
@ -348,9 +352,15 @@ async def run_async_fallback(
attempted.record(original_model_group)
for mg in fallback_model_group:
if mg == original_model_group:
target_model_group: Final = _get_fallback_target_model_group(mg)
canonical_target_model_group: Final = (
resolve_model_group_alias(alias_map, target_model_group) or target_model_group
if isinstance(target_model_group, str)
else None
)
if isinstance(mg, str) and canonical_target_model_group == canonical_original_model_group:
continue
if same_model_group_only and _get_fallback_target_model_group(mg) != original_model_group:
if same_model_group_only and canonical_target_model_group != canonical_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),
@ -367,13 +377,21 @@ async def run_async_fallback(
continue
attempted.record(attempt_key)
try:
# LOGGING
# Deployment exclusions belong to the fallback target that just failed.
# Keep them through that target's retries, then clear them only when
# advancing to a distinct trusted fallback target.
kwargs.pop("_excluded_deployment_ids", None)
kwargs = litellm_router.log_retry(kwargs=kwargs, e=original_exception)
verbose_router_logger.info("Falling back to model_group = %s", mask_sensitive_structure(mg))
if isinstance(mg, str):
kwargs["model"] = mg
elif isinstance(mg, dict):
kwargs.update(mg)
preserve_authenticated_team_context(
request_kwargs=kwargs,
authenticated_team_id=authenticated_team_id,
source_bucket=authenticated_team_bucket,
)
fallback_depth = fallback_depth + 1
kwargs[metadata_variable_name] = {
"original_model_group": original_model_group,
@ -393,7 +411,6 @@ async def run_async_fallback(
attempted_fallbacks=fallback_depth,
fallback_errors=(list(fallback_errors) if include_fallback_errors else None),
)
# callback for successfull_fallback_event():
await log_success_fallback_event(
original_model_group=original_model_group,
kwargs=kwargs,
@ -419,22 +436,7 @@ async def run_async_fallback(
async def log_success_fallback_event(original_model_group: str, kwargs: dict, original_exception: Exception):
"""
Log a successful fallback event to all registered callbacks.
Uses LoggingCallbackManager.get_custom_loggers_for_type() to get deduplicated
CustomLogger instances from all callback lists.
Args:
original_model_group (str): The original model group before fallback.
kwargs (dict): kwargs for the request
Note:
Errors during logging are caught and reported but do not interrupt the process.
"""
# Get deduplicated CustomLogger instances from all callback lists
custom_loggers: Final = litellm.logging_callback_manager.get_custom_loggers_for_type(CustomLogger)
for _callback_custom_logger in custom_loggers:
try:
await _callback_custom_logger.log_success_fallback_event(
@ -447,22 +449,7 @@ async def log_success_fallback_event(original_model_group: str, kwargs: dict, or
async def log_failure_fallback_event(original_model_group: str, kwargs: dict, original_exception: Exception):
"""
Log a failed fallback event to all registered callbacks.
Uses LoggingCallbackManager.get_custom_loggers_for_type() to get deduplicated
CustomLogger instances from all callback lists.
Args:
original_model_group (str): The original model group before fallback.
kwargs (dict): kwargs for the request
Note:
Errors during logging are caught and reported but do not interrupt the process.
"""
# Get deduplicated CustomLogger instances from all callback lists
custom_loggers: Final = litellm.logging_callback_manager.get_custom_loggers_for_type(CustomLogger)
for _callback_custom_logger in custom_loggers:
try:
await _callback_custom_logger.log_failure_fallback_event(
@ -474,30 +461,32 @@ async def log_failure_fallback_event(original_model_group: str, kwargs: dict, or
verbose_router_logger.error("Error in log_failure_fallback_event: %s", e)
def _is_non_standard_fallback_target(item: Any) -> bool:
if isinstance(item, str):
return True
if not isinstance(item, dict):
return False
return "model" in item
def _is_unambiguous_direct_fallback_dict(item: Any) -> bool:
if not isinstance(item, dict) or "model" not in item:
return False
model = item.get("model")
if not isinstance(model, list):
return True
return any(key != "model" and not isinstance(value, list) for key, value in item.items())
def _check_non_standard_fallback_format(fallbacks: list[Any] | None) -> bool:
"""
Checks if the fallbacks list is a list of strings or a list of dictionaries.
If
- List[str]: e.g. ["claude-3-haiku", "openai/o-1"]
- List[Dict[<LiteLLMParamsTypedDict>, Any]]: e.g. [{"model": "claude-3-haiku", "messages": [{"role": "user", "content": "Hey, how's it going?"}]}]
If [{"gpt-3.5-turbo": ["claude-3-haiku"]}] then standard format.
"""
"""Check whether ``fallbacks`` is a direct ordered list of fallback targets."""
if fallbacks is None or not isinstance(fallbacks, list) or len(fallbacks) == 0:
return False
if all(isinstance(item, str) for item in fallbacks):
if not all(_is_non_standard_fallback_target(item) for item in fallbacks):
return False
if any(isinstance(item, str) for item in fallbacks):
return True
elif all(isinstance(item, dict) for item in fallbacks):
for item in fallbacks:
for key in LiteLLMParamsTypedDict.__annotations__:
if key in item:
# If the value is a list, it's likely a standard fallback model group mapping
# (e.g. {"model": ["backup"]}) rather than a parameter override.
if not isinstance(item[key], list):
return True
return False
return any(_is_unambiguous_direct_fallback_dict(item) for item in fallbacks)
def run_non_standard_fallback_format(fallbacks: list[str] | list[dict[str, Any]], model_group: str):

View file

@ -0,0 +1,490 @@
from types import SimpleNamespace
from unittest.mock import AsyncMock, patch
import pytest
import litellm
from litellm.proxy.route_llm_request import _has_available_fallback, route_request
from litellm.router_utils.fallback_event_handlers import run_async_fallback
def _blocked_primary_router(
*,
fallback_litellm_params: dict | None = None,
**router_kwargs,
) -> litellm.Router:
fallback_params = {
"model": "openai/gpt-4o-mini",
"api_key": "sk-test",
"mock_response": "fallback response",
**(fallback_litellm_params or {}),
}
return litellm.Router(
model_list=[
{
"model_name": "primary-model",
"litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test"},
"model_info": {"blocked": True},
},
{
"model_name": "fallback-model",
"litellm_params": fallback_params,
},
],
model_group_alias={"public-model": "primary-model"},
fallbacks=[{"public-model": ["primary-model", "fallback-model"]}],
num_retries=0,
**router_kwargs,
)
@pytest.mark.asyncio
async def test_blocked_model_rejected_when_fallbacks_disabled_for_request():
router = _blocked_primary_router()
with patch("litellm.proxy.route_llm_request.mock_testing_params_allowed", return_value=False):
with pytest.raises(litellm.PermissionDeniedError, match="Model is blocked"):
await route_request(
data={
"model": "public-model",
"messages": [{"role": "user", "content": "Hello"}],
"disable_fallbacks": True,
},
llm_router=router,
user_model=None,
route_type="acompletion",
)
@pytest.mark.asyncio
async def test_blocked_model_evals_route_rejected_even_with_healthy_fallback():
router = _blocked_primary_router()
with (
patch("litellm.proxy.route_llm_request.mock_testing_params_allowed", return_value=False),
patch(
"litellm.proxy.route_llm_request._has_available_fallback",
new=AsyncMock(return_value=True),
) as fallback_gate,
):
with pytest.raises(litellm.PermissionDeniedError, match="Model is blocked"):
await route_request(
data={"model": "public-model"},
llm_router=router,
user_model=None,
route_type="alist_evals",
)
fallback_gate.assert_not_awaited()
@pytest.mark.asyncio
async def test_blocked_model_rejected_when_fallback_is_runtime_ineligible():
router = _blocked_primary_router(
fallback_litellm_params={"rpm": 0},
enable_pre_call_checks=True,
)
with patch("litellm.proxy.route_llm_request.mock_testing_params_allowed", return_value=False):
with pytest.raises(litellm.PermissionDeniedError, match="Model is blocked"):
await route_request(
data={
"model": "public-model",
"messages": [{"role": "user", "content": "Hello"}],
},
llm_router=router,
user_model=None,
route_type="acompletion",
)
class _RecordingHealthyRouter:
max_fallbacks = 5
def __init__(self) -> None:
self.fallbacks = [
{
"public-model": [
{
"model": "fallback-model",
"metadata": {"user_api_key_team_id": "attacker-team"},
"litellm_metadata": {"user_api_key_team_id": "attacker-team"},
}
]
}
]
self.request_kwargs = None
async def async_pre_routing_hook(self, **kwargs):
return None
async def async_get_healthy_deployments(self, *, request_kwargs, **kwargs):
self.request_kwargs = request_kwargs
return [{"model_info": {"id": "fallback-deployment"}}]
class _RoutingPluginFilteredRouter(_RecordingHealthyRouter):
routing_plugins = [object()]
def __init__(self) -> None:
super().__init__()
self.pre_routing_hook_called = False
async def async_pre_routing_hook(self, *, request_kwargs, **kwargs):
self.pre_routing_hook_called = True
request_kwargs["plugin_excluded_fallback"] = True
return None
async def async_get_healthy_deployments(self, *, request_kwargs, **kwargs):
self.request_kwargs = request_kwargs
if request_kwargs.get("plugin_excluded_fallback") is True:
return []
return [{"model_info": {"id": "fallback-deployment"}}]
@pytest.mark.asyncio
async def test_fallback_gate_runs_pre_routing_hook_before_health_check():
router = _RoutingPluginFilteredRouter()
assert not await _has_available_fallback(
llm_router=router,
model_name="public-model",
team_id=None,
request_data={"messages": [{"role": "user", "content": "Hello"}]},
)
assert router.pre_routing_hook_called is True
assert router.request_kwargs["plugin_excluded_fallback"] is True
class _StrategyRewriteRouter(_RecordingHealthyRouter):
routing_plugins = []
def __init__(self) -> None:
super().__init__()
self.pre_routing_hook_called = False
self.health_check_model = None
self.health_check_messages = None
async def async_pre_routing_hook(self, *, messages, **kwargs):
self.pre_routing_hook_called = True
rewritten_messages = [*messages, {"role": "system", "content": "strategy tier selected"}]
return SimpleNamespace(model="strategy-tier", messages=rewritten_messages, litellm_params=None)
async def async_get_healthy_deployments(self, *, model, messages, request_kwargs, **kwargs):
self.request_kwargs = request_kwargs
self.health_check_model = model
self.health_check_messages = messages
if model != "strategy-tier":
return []
return [{"model_info": {"id": "strategy-deployment"}}]
@pytest.mark.asyncio
async def test_fallback_gate_runs_strategy_pre_routing_hook_without_plugins():
router = _StrategyRewriteRouter()
assert await _has_available_fallback(
llm_router=router,
model_name="public-model",
team_id=None,
request_data={"messages": [{"role": "user", "content": "Hello"}]},
)
assert router.routing_plugins == []
assert router.pre_routing_hook_called is True
assert router.health_check_model == "strategy-tier"
assert router.health_check_messages[-1]["content"] == "strategy tier selected"
class _StrategyParamsRouter(_RecordingHealthyRouter):
routing_plugins = []
async def async_pre_routing_hook(self, *, messages, **kwargs):
return SimpleNamespace(
model="strategy-tier",
messages=messages,
litellm_params={"tags": ["strategy-only"], "temperature": 0.25},
)
async def async_get_healthy_deployments(self, *, model, request_kwargs, **kwargs):
self.request_kwargs = request_kwargs
if model != "strategy-tier" or request_kwargs.get("tags") != ["strategy-only"]:
return []
return [{"model_info": {"id": "strategy-deployment"}}]
@pytest.mark.asyncio
async def test_fallback_gate_applies_strategy_litellm_params_before_health_check():
router = _StrategyParamsRouter()
assert await _has_available_fallback(
llm_router=router,
model_name="public-model",
team_id=None,
request_data={"messages": [{"role": "user", "content": "Hello"}]},
)
assert router.request_kwargs["tags"] == ["strategy-only"]
assert router.request_kwargs["temperature"] == 0.25
@pytest.mark.asyncio
async def test_fallback_gate_skips_cross_model_fallback_for_provider_scoped_resource():
router = _RecordingHealthyRouter()
assert not await _has_available_fallback(
llm_router=router,
model_name="public-model",
team_id=None,
request_data={
"input_file_id": "file-provider-scoped",
"messages": [{"role": "user", "content": "Hello"}],
},
)
assert router.request_kwargs is None
@pytest.mark.asyncio
async def test_fallback_gate_allows_same_group_dict_fallback_through_model_alias_for_provider_scoped_resource():
router = _RecordingHealthyRouter()
router.model_group_alias = {"public-model": "primary-model"}
router.fallbacks = [
{
"public-model": [
{
"model": "primary-model",
"temperature": 0.2,
}
]
}
]
assert await _has_available_fallback(
llm_router=router,
model_name="public-model",
team_id=None,
request_data={
"input_file_id": "file-provider-scoped",
"messages": [{"role": "user", "content": "Hello"}],
},
)
assert router.request_kwargs["model"] == "primary-model"
assert router.request_kwargs["temperature"] == 0.2
@pytest.mark.asyncio
async def test_fallback_gate_keeps_authenticated_team_authoritative():
router = _RecordingHealthyRouter()
assert await _has_available_fallback(
llm_router=router,
model_name="public-model",
team_id="trusted-team",
request_data={
"metadata": {"user_api_key_team_id": "trusted-team"},
"messages": [{"role": "user", "content": "Hello"}],
},
)
assert router.request_kwargs["metadata"]["user_api_key_team_id"] == "trusted-team"
assert router.request_kwargs["litellm_metadata"]["user_api_key_team_id"] == "trusted-team"
@pytest.mark.asyncio
async def test_fallback_gate_rejects_specific_deployment_from_another_team():
router = litellm.Router(
model_list=[
{
"model_name": "primary-model",
"litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test"},
"model_info": {"blocked": True},
},
{
"model_name": "fallback-model",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "sk-test",
"mock_response": "team-b fallback",
},
"model_info": {"id": "team-b-deployment", "team_id": "team-b"},
},
],
model_group_alias={"public-model": "primary-model"},
fallbacks=[{"public-model": ["team-b-deployment"]}],
num_retries=0,
)
request_data = {
"metadata": {"user_api_key_team_id": "team-a"},
"messages": [{"role": "user", "content": "Hello"}],
}
assert not await _has_available_fallback(
llm_router=router,
model_name="public-model",
team_id="team-a",
request_data=request_data,
)
request_data["metadata"]["user_api_key_team_id"] = "team-b"
assert await _has_available_fallback(
llm_router=router,
model_name="public-model",
team_id="team-b",
request_data=request_data,
)
@pytest.mark.asyncio
async def test_fallback_gate_accepts_supported_list_of_dict_fallback_format():
router = _RecordingHealthyRouter()
fallback_messages = [{"role": "user", "content": "Use the fallback prompt"}]
router.fallbacks = [
{
"model": "fallback-model",
"messages": fallback_messages,
}
]
assert await _has_available_fallback(
llm_router=router,
model_name="public-model",
team_id=None,
request_data={"messages": [{"role": "user", "content": "Hello"}]},
)
assert router.request_kwargs["model"] == "fallback-model"
assert router.request_kwargs["messages"] == fallback_messages
@pytest.mark.asyncio
async def test_blocked_model_uses_server_fallback_instead_of_request_supplied_fallback():
router = litellm.Router(
model_list=[
{
"model_name": "primary-model",
"litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test"},
"model_info": {"blocked": True},
},
{
"model_name": "fallback-model",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "sk-test",
"mock_response": "server fallback response",
},
},
{
"model_name": "restricted-model",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "sk-test",
"mock_response": "request fallback response",
},
},
],
model_group_alias={"public-model": "primary-model"},
fallbacks=[{"public-model": ["fallback-model"]}],
num_retries=0,
)
with patch("litellm.proxy.route_llm_request.mock_testing_params_allowed", return_value=False):
response = await route_request(
data={
"model": "public-model",
"messages": [{"role": "user", "content": "Hello"}],
"fallbacks": [{"model": "restricted-model"}],
},
llm_router=router,
user_model=None,
route_type="acompletion",
)
assert response.choices[0].message.content == "server fallback response"
class _RecordingFallbackRouter:
def __init__(self) -> None:
self.request_kwargs = None
def log_retry(self, kwargs, e):
return kwargs
async def async_function_with_fallbacks(self, *args, **kwargs):
self.request_kwargs = kwargs
return "fallback response"
@pytest.mark.asyncio
async def test_runtime_fallback_keeps_authenticated_team_authoritative():
router = _RecordingFallbackRouter()
async def _original_function():
return None
with (
patch(
"litellm.router_utils.fallback_event_handlers.add_fallback_headers_to_response",
side_effect=lambda response, **kwargs: response,
),
patch(
"litellm.router_utils.fallback_event_handlers.log_success_fallback_event",
new=AsyncMock(),
),
):
response = await run_async_fallback(
litellm_router=router,
fallback_model_group=[
{
"model": "fallback-model",
"metadata": {"user_api_key_team_id": "attacker-team"},
"litellm_metadata": {"user_api_key_team_id": "attacker-team"},
}
],
original_model_group="primary-model",
original_exception=RuntimeError("primary failed"),
max_fallbacks=1,
fallback_depth=0,
original_function=_original_function,
metadata={"user_api_key_team_id": "trusted-team"},
)
assert response == "fallback response"
assert router.request_kwargs["metadata"]["user_api_key_team_id"] == "trusted-team"
assert router.request_kwargs["litellm_metadata"]["user_api_key_team_id"] == "trusted-team"
@pytest.mark.asyncio
async def test_runtime_provider_scoped_alias_allows_same_group_dict_fallback():
router = _RecordingFallbackRouter()
router.model_group_alias = {"public-model": "primary-model"}
async def _original_function():
return None
with (
patch(
"litellm.router_utils.fallback_event_handlers.add_fallback_headers_to_response",
side_effect=lambda response, **kwargs: response,
),
patch(
"litellm.router_utils.fallback_event_handlers.log_success_fallback_event",
new=AsyncMock(),
),
):
response = await run_async_fallback(
litellm_router=router,
fallback_model_group=[{"model": "primary-model", "temperature": 0.2}],
original_model_group="public-model",
original_exception=RuntimeError("primary failed"),
max_fallbacks=1,
fallback_depth=0,
original_function=_original_function,
input_file_id="file-provider-scoped",
)
assert response == "fallback response"
assert router.request_kwargs["model"] == "primary-model"
assert router.request_kwargs["temperature"] == 0.2

View file

@ -0,0 +1,321 @@
from types import SimpleNamespace
from unittest.mock import patch
import pytest
import litellm
from litellm.proxy.route_llm_request import route_request
def _stateful_rewrite_router() -> litellm.Router:
return litellm.Router(
model_list=[
{
"model_name": "primary-model",
"litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test"},
"model_info": {"blocked": True},
},
{
"model_name": "fallback-model",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "sk-test",
"mock_response": "first fallback response",
},
},
{
"model_name": "safe-fallback-model",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "sk-test",
"mock_response": "safe fallback response",
},
},
{
"model_name": "restricted-model",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "sk-test",
"mock_response": "restricted response",
},
},
],
model_group_alias={"public-model": "primary-model"},
fallbacks=[{"public-model": ["fallback-model", "safe-fallback-model"]}],
num_retries=0,
)
@pytest.mark.asyncio
async def test_blocked_primary_cannot_be_rewritten_before_trusted_fallback():
router = litellm.Router(
model_list=[
{
"model_name": "primary-model",
"litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test"},
"model_info": {"blocked": True},
},
{
"model_name": "fallback-model",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "sk-test",
"mock_response": "server fallback response",
},
},
{
"model_name": "restricted-model",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "sk-test",
"mock_response": "rewritten primary response",
},
},
],
model_group_alias={"public-model": "primary-model"},
fallbacks=[{"public-model": ["fallback-model"]}],
num_retries=0,
)
async def rewrite_only_blocked_primary(*, model, messages, **kwargs):
if model != "public-model":
return None
return SimpleNamespace(
model="restricted-model",
messages=messages,
litellm_params=None,
)
router.async_pre_routing_hook = rewrite_only_blocked_primary
with patch("litellm.proxy.route_llm_request.mock_testing_params_allowed", return_value=False):
response = await route_request(
data={
"model": "public-model",
"messages": [{"role": "user", "content": "Hello"}],
},
llm_router=router,
user_model=None,
route_type="acompletion",
)
assert response.choices[0].message.content == "server fallback response"
@pytest.mark.asyncio
async def test_stateful_fallback_rewrite_cannot_escape_preflight_validation():
router = _stateful_rewrite_router()
fallback_hook_calls = 0
async def stateful_fallback_rewrite(*, model, messages, **kwargs):
nonlocal fallback_hook_calls
if model != "fallback-model":
return None
fallback_hook_calls += 1
if fallback_hook_calls == 1:
return None
return SimpleNamespace(
model="restricted-model",
messages=messages,
litellm_params=None,
)
router.async_pre_routing_hook = stateful_fallback_rewrite
with patch("litellm.proxy.route_llm_request.mock_testing_params_allowed", return_value=False):
response = await route_request(
data={
"model": "public-model",
"messages": [{"role": "user", "content": "Hello"}],
},
llm_router=router,
user_model=None,
route_type="acompletion",
)
assert fallback_hook_calls == 2
assert response.choices[0].message.content == "safe fallback response"
@pytest.mark.asyncio
async def test_stateful_fallback_concrete_deployment_cannot_escape_preflight_validation():
router = _stateful_rewrite_router()
restricted_deployment_id = next(
deployment["model_info"]["id"]
for deployment in router.model_list
if deployment["model_name"] == "restricted-model"
)
fallback_hook_calls = 0
async def stateful_fallback_rewrite(*, model, messages, **kwargs):
nonlocal fallback_hook_calls
if model != "fallback-model":
return None
fallback_hook_calls += 1
if fallback_hook_calls == 1:
return None
return SimpleNamespace(
model=restricted_deployment_id,
messages=messages,
litellm_params=None,
)
router.async_pre_routing_hook = stateful_fallback_rewrite
with patch("litellm.proxy.route_llm_request.mock_testing_params_allowed", return_value=False):
response = await route_request(
data={
"model": "public-model",
"messages": [{"role": "user", "content": "Hello"}],
},
llm_router=router,
user_model=None,
route_type="acompletion",
)
assert fallback_hook_calls == 2
assert response.choices[0].message.content == "safe fallback response"
@pytest.mark.asyncio
async def test_failed_concrete_fallback_consumes_exclusions_before_next_trusted_target():
router = litellm.Router(
model_list=[
{
"model_name": "primary-model",
"litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test"},
"model_info": {"blocked": True},
},
{
"model_name": "fallback-model",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "sk-test",
"mock_response": Exception("first fallback failed"),
},
},
{
"model_name": "safe-fallback-model",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "sk-test",
"mock_response": "safe fallback response",
},
},
],
model_group_alias={"public-model": "primary-model"},
fallbacks=[{"public-model": ["fallback-model", "safe-fallback-model"]}],
num_retries=0,
)
fallback_deployment_id = next(
deployment["model_info"]["id"]
for deployment in router.model_list
if deployment["model_name"] == "fallback-model"
)
fallback_hook_calls = 0
async def select_concrete_first_fallback(*, model, messages, **kwargs):
nonlocal fallback_hook_calls
if model != "fallback-model":
return None
fallback_hook_calls += 1
return SimpleNamespace(
model=fallback_deployment_id,
messages=messages,
litellm_params=None,
)
router.async_pre_routing_hook = select_concrete_first_fallback
with patch("litellm.proxy.route_llm_request.mock_testing_params_allowed", return_value=False):
response = await route_request(
data={
"model": "public-model",
"messages": [{"role": "user", "content": "Hello"}],
},
llm_router=router,
user_model=None,
route_type="acompletion",
)
assert fallback_hook_calls == 2
assert response.choices[0].message.content == "safe fallback response"
@pytest.mark.asyncio
async def test_retry_cannot_escape_preflight_validated_concrete_deployment():
router = litellm.Router(
model_list=[
{
"model_name": "primary-model",
"litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test"},
"model_info": {"blocked": True},
},
{
"model_name": "fallback-model",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "sk-test",
"mock_response": Exception("validated fallback failed"),
},
"model_info": {"id": "validated-fallback-deployment"},
},
{
"model_name": "fallback-model",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "sk-test",
"mock_response": "unvalidated retry response",
},
"model_info": {"id": "unvalidated-fallback-deployment"},
},
{
"model_name": "safe-fallback-model",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "sk-test",
"mock_response": "safe fallback response",
},
},
],
model_group_alias={"public-model": "primary-model"},
fallbacks=[{"public-model": ["fallback-model", "safe-fallback-model"]}],
num_retries=1,
)
hook_calls = 0
async def pin_then_attempt_retry_escape(*, model, messages, **kwargs):
nonlocal hook_calls
if model not in {
"fallback-model",
"validated-fallback-deployment",
"unvalidated-fallback-deployment",
}:
return None
hook_calls += 1
target = (
"unvalidated-fallback-deployment"
if hook_calls >= 3
else "validated-fallback-deployment"
)
return SimpleNamespace(
model=target,
messages=messages,
litellm_params=None,
)
router.async_pre_routing_hook = pin_then_attempt_retry_escape
with patch("litellm.proxy.route_llm_request.mock_testing_params_allowed", return_value=False):
response = await route_request(
data={
"model": "public-model",
"messages": [{"role": "user", "content": "Hello"}],
},
llm_router=router,
user_model=None,
route_type="acompletion",
)
assert hook_calls >= 3
assert response.choices[0].message.content == "safe fallback response"

View file

@ -0,0 +1,102 @@
from unittest.mock import AsyncMock
import pytest
from litellm.proxy.route_llm_request import _get_available_fallback_request
class _ListFallbackRouter:
max_fallbacks = 5
model_group_alias = {}
def __init__(self, fallbacks, healthy_models):
self.fallbacks = fallbacks
self.healthy_models = set(healthy_models)
self.model_list = [
{"model_info": {"id": "primary-deployment"}},
{"model_info": {"id": "first-deployment"}},
{"model_info": {"id": "second-deployment"}},
{"model_info": {"id": "later-deployment"}},
]
self.async_pre_routing_hook = AsyncMock(return_value=None)
async def async_get_healthy_deployments(self, *, model, **kwargs):
if model not in self.healthy_models:
return []
deployment_id = {
"first-model": "first-deployment",
"second-model": "second-deployment",
"later-model": "later-deployment",
}[model]
return [{"model_info": {"id": deployment_id}}]
@pytest.mark.asyncio
async def test_list_valued_direct_fallback_uses_first_healthy_candidate():
router = _ListFallbackRouter(
fallbacks=[{"model": ["first-model", "second-model"], "temperature": 0.2}],
healthy_models={"second-model"},
)
request = await _get_available_fallback_request(
llm_router=router,
model_name="blocked-model",
team_id=None,
request_data={"messages": [{"role": "user", "content": "hello"}]},
)
assert request is not None
assert request["model"] == "second-model"
assert request["temperature"] == 0.2
assert request["fallbacks"] == []
assert [call.kwargs["model"] for call in router.async_pre_routing_hook.await_args_list] == [
"first-model",
"second-model",
]
@pytest.mark.asyncio
async def test_list_valued_direct_fallback_preserves_remaining_candidates_in_trusted_tail():
router = _ListFallbackRouter(
fallbacks=[
{"model": ["first-model", "second-model"], "temperature": 0.2},
"later-model",
],
healthy_models={"first-model", "second-model", "later-model"},
)
request = await _get_available_fallback_request(
llm_router=router,
model_name="blocked-model",
team_id=None,
request_data={"messages": [{"role": "user", "content": "hello"}]},
)
assert request is not None
assert request["model"] == "first-model"
assert request["fallbacks"] == [
{"model": "second-model", "temperature": 0.2},
"later-model",
]
@pytest.mark.asyncio
async def test_later_list_valued_fallback_is_expanded_in_trusted_tail():
router = _ListFallbackRouter(
fallbacks=[
{"model": ["first-model"], "temperature": 0.2},
{"model": ["later-model"]},
],
healthy_models={"first-model", "later-model"},
)
request = await _get_available_fallback_request(
llm_router=router,
model_name="blocked-model",
team_id=None,
request_data={"messages": [{"role": "user", "content": "hello"}]},
)
assert request is not None
assert request["model"] == "first-model"
assert request["fallbacks"] == [{"model": "later-model"}]

View file

@ -0,0 +1,157 @@
from unittest.mock import AsyncMock, patch
import pytest
from litellm.proxy.route_llm_request import _get_available_fallback_request
from litellm.router_utils.common_utils import filter_team_based_models
from litellm.router_utils.fallback_event_handlers import run_async_fallback
class _PreflightRouter:
max_fallbacks = 5
model_group_alias = None
def __init__(self) -> None:
self.fallbacks = [
{
"public-model": [
{
"model": "fallback-model",
"metadata": {"user_api_key_auth": "attacker-auth"},
}
]
}
]
self.model_list = [
{
"model_name": "fallback-model",
"model_info": {"id": "fallback-deployment"},
}
]
self.request_kwargs = None
async def async_pre_routing_hook(self, **kwargs):
return None
async def async_get_healthy_deployments(self, *, request_kwargs, **kwargs):
self.request_kwargs = request_kwargs
return [{"model_info": {"id": "fallback-deployment"}}]
@pytest.mark.asyncio
async def test_fallback_preflight_preserves_authenticated_api_key_context():
router = _PreflightRouter()
trusted_auth = object()
fallback_request = await _get_available_fallback_request(
llm_router=router,
model_name="public-model",
team_id=None,
request_data={
"metadata": {"user_api_key_auth": trusted_auth},
"messages": [{"role": "user", "content": "Hello"}],
},
)
assert fallback_request is not None
assert router.request_kwargs["metadata"]["user_api_key_auth"] is trusted_auth
assert fallback_request["metadata"]["user_api_key_auth"] is trusted_auth
@pytest.mark.asyncio
async def test_fallback_preflight_strips_injected_api_key_context_when_none_authenticated():
router = _PreflightRouter()
fallback_request = await _get_available_fallback_request(
llm_router=router,
model_name="public-model",
team_id=None,
request_data={"messages": [{"role": "user", "content": "Hello"}]},
)
assert fallback_request is not None
assert "user_api_key_auth" not in router.request_kwargs.get("metadata", {})
assert "user_api_key_auth" not in fallback_request.get("metadata", {})
def test_retry_scoped_deployment_exclusions_survive_router_pop_until_fallback_advances():
deployments = [
{"model_info": {"id": "validated-deployment"}},
{"model_info": {"id": "unvalidated-deployment"}},
]
request_kwargs = {
"fallback_depth": 1,
"_excluded_deployment_ids": ["unvalidated-deployment"],
}
first_lookup = filter_team_based_models(deployments, request_kwargs)
assert [item["model_info"]["id"] for item in first_lookup] == ["validated-deployment"]
# Router health selection consumes the public exclusion key after team
# filtering. A retry at the same fallback depth must restore it.
request_kwargs.pop("_excluded_deployment_ids")
retry_lookup = filter_team_based_models(deployments, request_kwargs)
assert [item["model_info"]["id"] for item in retry_lookup] == ["validated-deployment"]
assert request_kwargs["_excluded_deployment_ids"] == ["unvalidated-deployment"]
# Moving to the next trusted fallback increments fallback_depth. The old
# target's exclusions must no longer constrain that new fallback.
request_kwargs.pop("_excluded_deployment_ids")
request_kwargs["fallback_depth"] = 2
next_fallback_lookup = filter_team_based_models(deployments, request_kwargs)
assert [item["model_info"]["id"] for item in next_fallback_lookup] == [
"validated-deployment",
"unvalidated-deployment",
]
class _RuntimeFallbackRouter:
model_group_alias = None
def __init__(self) -> None:
self.request_kwargs = None
def log_retry(self, kwargs, e):
return kwargs
async def async_function_with_fallbacks(self, *args, **kwargs):
self.request_kwargs = kwargs
return "fallback response"
@pytest.mark.asyncio
async def test_runtime_fallback_preserves_authenticated_api_key_context():
router = _RuntimeFallbackRouter()
trusted_auth = object()
async def _original_function():
return None
with (
patch(
"litellm.router_utils.fallback_event_handlers.add_fallback_headers_to_response",
side_effect=lambda response, **kwargs: response,
),
patch(
"litellm.router_utils.fallback_event_handlers.log_success_fallback_event",
new=AsyncMock(),
),
):
response = await run_async_fallback(
litellm_router=router,
fallback_model_group=[
{
"model": "fallback-model",
"metadata": {"user_api_key_auth": "attacker-auth"},
}
],
original_model_group="primary-model",
original_exception=RuntimeError("primary failed"),
max_fallbacks=1,
fallback_depth=0,
original_function=_original_function,
metadata={"user_api_key_auth": trusted_auth},
)
assert response == "fallback response"
assert router.request_kwargs["metadata"]["user_api_key_auth"] is trusted_auth

View file

@ -0,0 +1,42 @@
from litellm.router_utils.fallback_event_handlers import (
_check_non_standard_fallback_format,
)
def test_mixed_direct_fallback_targets_are_non_standard_format() -> None:
fallbacks = [
"backup-a",
{"model": "backup-b", "temperature": 0},
]
assert _check_non_standard_fallback_format(fallbacks) is True
def test_model_group_mapping_remains_standard_format() -> None:
fallbacks = [{"primary": ["backup-a", {"model": "backup-b"}]}]
assert _check_non_standard_fallback_format(fallbacks) is False
def test_model_key_mapping_with_list_targets_remains_standard_format() -> None:
fallbacks = [{"model": ["backup-a", "backup-b"]}]
assert _check_non_standard_fallback_format(fallbacks) is False
def test_multi_key_list_mapping_remains_standard_format() -> None:
fallbacks = [{"model": ["qwen-backup"], "region": ["us-east-1"]}]
assert _check_non_standard_fallback_format(fallbacks) is False
def test_list_valued_direct_model_with_request_overrides_is_non_standard_format() -> None:
fallbacks = [{"model": ["backup-a", "backup-b"], "temperature": 0}]
assert _check_non_standard_fallback_format(fallbacks) is True
def test_mixed_string_and_list_valued_model_target_is_non_standard_format() -> None:
fallbacks = ["backup-a", {"model": ["backup-b", "backup-c"]}]
assert _check_non_standard_fallback_format(fallbacks) is True