mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge 1f2b7bdce3 into 3930c5bab6
This commit is contained in:
commit
bface5024e
3 changed files with 276 additions and 0 deletions
|
|
@ -693,3 +693,77 @@ def _get_tags_from_request_kwargs(
|
|||
typed_litellm_params: Final[Mapping[str, object]] = litellm_params
|
||||
return _tags_in_metadata(typed_litellm_params.get(resolved_variable_name))
|
||||
return []
|
||||
|
||||
|
||||
def can_satisfy_confirmed_routing_tags(
|
||||
llm_router_instance: LitellmRouter,
|
||||
model: str,
|
||||
request_kwargs: Mapping[str, object] | None = None,
|
||||
metadata_variable_name: Literal["metadata", "litellm_metadata"] | None = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Whether attempting ``model`` can possibly succeed under the request's
|
||||
``tag_routing_prefix``-confirmed tags.
|
||||
|
||||
Returns False only when confirmed routing tags make success structurally
|
||||
impossible for every deployment in the model group — so
|
||||
``run_async_fallback`` can skip that leg without raising, logging ERROR, or
|
||||
counting a deployment failure. Returns True when the check does not apply
|
||||
(no prefix, no confirmed tags, tag filtering off) or the group might still
|
||||
serve the request.
|
||||
"""
|
||||
routing_prefix: Final = getattr(llm_router_instance, "tag_routing_prefix", None) or ""
|
||||
if not routing_prefix or request_kwargs is None:
|
||||
return True
|
||||
|
||||
request_tags: Final = _get_tags_from_request_kwargs(request_kwargs, metadata_variable_name)
|
||||
rewritten_tags, routing_confirmed = _strip_routing_prefix(request_tags, routing_prefix)
|
||||
if not routing_confirmed:
|
||||
return True
|
||||
|
||||
try:
|
||||
deployments: Final = llm_router_instance._get_all_deployments(model_name=model)
|
||||
except Exception: # noqa: BLE001 # fail safe toward attempting the leg on lookup errors
|
||||
return True
|
||||
if not deployments:
|
||||
return True
|
||||
|
||||
request_enable_tag_filtering: Final = request_kwargs.get("enable_tag_filtering")
|
||||
chain_enable_tag_filtering: Final = _chain_tag_filtering_override(llm_router_instance, model, deployments)
|
||||
router_enable_tag_filtering: Final = getattr(llm_router_instance, "enable_tag_filtering", False)
|
||||
chain_default: Final = (
|
||||
chain_enable_tag_filtering if chain_enable_tag_filtering is not None else router_enable_tag_filtering
|
||||
)
|
||||
if request_enable_tag_filtering is not True and chain_default is not True:
|
||||
return True
|
||||
|
||||
required_tags, positive_tags, excluded_patterns = _split_tags(rewritten_tags)
|
||||
excluded_set: Final = frozenset(excluded_patterns)
|
||||
required_set: Final = frozenset(required_tags)
|
||||
candidates: Final = _require_all_tags(_exclude_deployments(deployments, excluded_set), required_set)
|
||||
|
||||
confirmed_required: Final = required_set & routing_confirmed
|
||||
confirmed_excluded: Final = excluded_set & routing_confirmed
|
||||
if not candidates and (confirmed_required or confirmed_excluded):
|
||||
return False
|
||||
|
||||
confirmed_positive: Final = [tag for tag in positive_tags if tag in routing_confirmed]
|
||||
if not confirmed_positive:
|
||||
return True
|
||||
|
||||
# A tag_regex deployment may still match via request headers even when plain
|
||||
# tags do not; leave those legs to the normal attempt path.
|
||||
if any(d.get("litellm_params", MappingProxyType({})).get("tag_regex") for d in candidates):
|
||||
return True
|
||||
|
||||
# Confirmed positive tags force a hard deny when nothing matches (defaults /
|
||||
# fail-open do not apply). Skip the leg rather than attempting it.
|
||||
match_any: Final = getattr(llm_router_instance, "tag_filtering_match_any", True)
|
||||
return any(
|
||||
is_valid_deployment_tag(
|
||||
d.get("litellm_params", MappingProxyType({})).get("tags") or [],
|
||||
positive_tags,
|
||||
match_any,
|
||||
)
|
||||
for d in candidates
|
||||
)
|
||||
|
|
|
|||
|
|
@ -545,6 +545,30 @@ async def _is_fallback_target_within_budget(
|
|||
return False
|
||||
|
||||
|
||||
def _is_fallback_target_tag_satisfiable(
|
||||
litellm_router: LitellmRouter,
|
||||
fallback_entry: str | Mapping[str, object],
|
||||
kwargs: Mapping[str, object],
|
||||
) -> bool:
|
||||
"""
|
||||
Skip a fallback leg whose model group cannot satisfy confirmed tag-routing
|
||||
tags. Mirrors the authorized/budget pre-checks: structurally unsatisfiable
|
||||
legs must not be attempted, raised, logged at ERROR, or counted as failures.
|
||||
"""
|
||||
target: Final = _get_fallback_target_model_group(fallback_entry)
|
||||
if target is None:
|
||||
return True
|
||||
from litellm.router_strategy.tag_based_routing import can_satisfy_confirmed_routing_tags
|
||||
|
||||
if can_satisfy_confirmed_routing_tags(litellm_router, target, kwargs):
|
||||
return True
|
||||
verbose_router_logger.info(
|
||||
"Skipping fallback to model_group = %s: no deployment can satisfy confirmed tag routing tags",
|
||||
mask_sensitive_structure(fallback_entry),
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
def references_provider_scoped_resource(kwargs: Mapping[str, object]) -> bool:
|
||||
"""
|
||||
True when a file, batch, or fine-tuning job operation names an id that only exists
|
||||
|
|
@ -654,6 +678,8 @@ async def run_async_fallback(
|
|||
continue
|
||||
if not await _is_fallback_target_within_budget(litellm_router, mg, original_model_group, kwargs):
|
||||
continue
|
||||
if not _is_fallback_target_tag_satisfiable(litellm_router, mg, kwargs):
|
||||
continue
|
||||
attempt_key = fallback_attempt_key(mg)
|
||||
if attempt_key is not None:
|
||||
if attempt_key in attempted:
|
||||
|
|
|
|||
|
|
@ -547,6 +547,182 @@ async def test_run_async_fallback_does_not_consult_access_check_for_same_model_g
|
|||
assert router.access_checks == []
|
||||
|
||||
|
||||
|
||||
|
||||
class TagRoutingFallbackRouter:
|
||||
"""Records fallback attempts while exposing tag-routing deployment lookups."""
|
||||
|
||||
fallback_access_check = None
|
||||
fallback_budget_check = None
|
||||
enable_tag_filtering = True
|
||||
tag_routing_prefix = "route:"
|
||||
tag_filtering_match_any = True
|
||||
|
||||
def __init__(self, deployments_by_group: dict[str, list[dict]]):
|
||||
self.deployments_by_group = deployments_by_group
|
||||
self.attempted_model_groups = []
|
||||
|
||||
def log_retry(self, kwargs, e):
|
||||
return kwargs
|
||||
|
||||
def _get_all_deployments(self, model_name: str, **kwargs):
|
||||
return list(self.deployments_by_group.get(model_name, []))
|
||||
|
||||
async def async_function_with_fallbacks(self, *args, **kwargs):
|
||||
self.attempted_model_groups.append(kwargs.get("model"))
|
||||
return StreamingWrapper()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_async_fallback_skips_legs_unsatisfiable_under_confirmed_tags(caplog):
|
||||
"""Confirmed tag_routing_prefix tags must skip structurally impossible legs at INFO,
|
||||
not attempt them (which would raise/log ERROR and inflate failure metrics)."""
|
||||
import logging
|
||||
|
||||
router = TagRoutingFallbackRouter(
|
||||
{
|
||||
"chain/2-vertex": [
|
||||
{
|
||||
"model_name": "chain/2-vertex",
|
||||
"litellm_params": {"tags": ["inference:vertex", "default"]},
|
||||
}
|
||||
],
|
||||
"chain/3-bedrock": [
|
||||
{
|
||||
"model_name": "chain/3-bedrock",
|
||||
"litellm_params": {"tags": ["inference:bedrock", "default"]},
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
with caplog.at_level(logging.INFO, logger="LiteLLM Router"):
|
||||
response = await run_async_fallback(
|
||||
litellm_router=router,
|
||||
fallback_model_group=["chain/2-vertex", "chain/3-bedrock"],
|
||||
original_model_group="chain/1-openai",
|
||||
original_exception=RuntimeError("primary failed"),
|
||||
max_fallbacks=5,
|
||||
fallback_depth=0,
|
||||
model="chain/1-openai",
|
||||
metadata={"tags": ["route:inference:bedrock"]},
|
||||
)
|
||||
|
||||
assert router.attempted_model_groups == ["chain/3-bedrock"]
|
||||
assert response._hidden_params["additional_headers"]["x-litellm-attempted-fallbacks"] == 1
|
||||
skip_logs = [r for r in caplog.records if "no deployment can satisfy confirmed tag routing tags" in r.getMessage()]
|
||||
assert len(skip_logs) == 1
|
||||
assert skip_logs[0].levelno == logging.INFO
|
||||
assert "chain/2-vertex" in skip_logs[0].getMessage()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_async_fallback_raises_original_when_all_legs_tag_unsatisfiable():
|
||||
router = TagRoutingFallbackRouter(
|
||||
{
|
||||
"chain/2-vertex": [
|
||||
{
|
||||
"model_name": "chain/2-vertex",
|
||||
"litellm_params": {"tags": ["inference:vertex", "default"]},
|
||||
}
|
||||
],
|
||||
"chain/3-anthropic": [
|
||||
{
|
||||
"model_name": "chain/3-anthropic",
|
||||
"litellm_params": {"tags": ["inference:anthropic", "default"]},
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError, match="primary failed"):
|
||||
await run_async_fallback(
|
||||
litellm_router=router,
|
||||
fallback_model_group=["chain/2-vertex", "chain/3-anthropic"],
|
||||
original_model_group="chain/1-openai",
|
||||
original_exception=RuntimeError("primary failed"),
|
||||
max_fallbacks=5,
|
||||
fallback_depth=0,
|
||||
model="chain/1-openai",
|
||||
metadata={"tags": ["route:inference:bedrock"]},
|
||||
)
|
||||
|
||||
assert router.attempted_model_groups == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_async_fallback_does_not_skip_when_tag_routing_prefix_unset():
|
||||
router = TagRoutingFallbackRouter(
|
||||
{
|
||||
"chain/2-vertex": [
|
||||
{
|
||||
"model_name": "chain/2-vertex",
|
||||
"litellm_params": {"tags": ["inference:vertex", "default"]},
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
router.tag_routing_prefix = ""
|
||||
|
||||
await run_async_fallback(
|
||||
litellm_router=router,
|
||||
fallback_model_group=["chain/2-vertex"],
|
||||
original_model_group="chain/1-openai",
|
||||
original_exception=RuntimeError("primary failed"),
|
||||
max_fallbacks=3,
|
||||
fallback_depth=0,
|
||||
model="chain/1-openai",
|
||||
metadata={"tags": ["route:inference:bedrock"]},
|
||||
)
|
||||
|
||||
assert router.attempted_model_groups == ["chain/2-vertex"]
|
||||
|
||||
|
||||
def test_can_satisfy_confirmed_routing_tags_false_for_mismatched_single_deployment():
|
||||
from litellm.router_strategy.tag_based_routing import can_satisfy_confirmed_routing_tags
|
||||
|
||||
router = TagRoutingFallbackRouter(
|
||||
{
|
||||
"chain/2-vertex": [
|
||||
{
|
||||
"model_name": "chain/2-vertex",
|
||||
"litellm_params": {"tags": ["inference:vertex", "default"]},
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
assert (
|
||||
can_satisfy_confirmed_routing_tags(
|
||||
router,
|
||||
"chain/2-vertex",
|
||||
{"metadata": {"tags": ["route:inference:bedrock"]}},
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
|
||||
def test_can_satisfy_confirmed_routing_tags_true_when_deployment_matches():
|
||||
from litellm.router_strategy.tag_based_routing import can_satisfy_confirmed_routing_tags
|
||||
|
||||
router = TagRoutingFallbackRouter(
|
||||
{
|
||||
"chain/3-bedrock": [
|
||||
{
|
||||
"model_name": "chain/3-bedrock",
|
||||
"litellm_params": {"tags": ["inference:bedrock", "default"]},
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
assert (
|
||||
can_satisfy_confirmed_routing_tags(
|
||||
router,
|
||||
"chain/3-bedrock",
|
||||
{"metadata": {"tags": ["route:inference:bedrock"]}},
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
class RecordingFailRouter:
|
||||
fallback_access_check = None
|
||||
fallback_budget_check = None
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue