Merge pull request #36627 from BerriAI/litellm_fix_autorouter_untagged_hijack
Some checks failed
Unit Tests: MCP, Secrets, Containers & Misc / misc (push) Waiting to run
Unit Tests: Proxy Auth & Key Management / proxy-auth (push) Waiting to run
Unit Tests: Proxy DB Operations / assert-shard-coverage (push) Waiting to run
Unit Tests: Proxy DB Operations / auth-checks (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / budgets (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / custom-logging (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / db-and-spend (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / endpoints-and-responses (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / guardrails-hooks (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / jwt-and-keys (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / key-generation (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / logging-misc (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-runtime (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-server-core (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-utils (push) Blocked by required conditions
Unit Tests: Proxy API Endpoints / proxy-endpoints (push) Waiting to run
Unit Tests: Proxy API Endpoints / proxy-server (push) Waiting to run
Unit Tests: Proxy Infrastructure / proxy-infra (push) Waiting to run
Unit Tests: Proxy Legacy Tests / auth-and-jwt (push) Waiting to run
Unit Tests: Proxy Legacy Tests / key-generation (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-config (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-response-and-misc (push) Waiting to run
CI Coverage / assert-ci-coverage (push) Waiting to run
CodSpeed Benchmarks / benchmarks (push) Waiting to run
Publish basedpyright base counts / publish (push) Waiting to run
Code Quality Checks / code-quality (push) Waiting to run
UI Unit Tests / ui-unit-tests (push) Waiting to run
Unit Tests: Core Utilities / core-utils (push) Waiting to run
Unit Tests: Documentation Validation / documentation (push) Waiting to run
Unit Tests: Enterprise, Google GenAI & Routing / enterprise-routing (push) Waiting to run
Unit Tests: Integrations (Callbacks & Logging) / integrations (push) Waiting to run
Unit Tests: LLM Provider Transformations / All Other Providers (push) Waiting to run
Unit Tests: LLM Provider Transformations / Vertex AI (push) Waiting to run
GitHub Actions Security Analysis / zizmor (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-server (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-server-extras (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-token-counter (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-user-auth-and-spend (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-utils (push) Waiting to run
Unit Tests: Responses, Caching & Types / responses-caching-types (push) Waiting to run
Terraform Modules / fmt, validate, test (aws) (push) Has been cancelled

fix(router): let untagged requests bypass a tagged pre-routing strategy on shared model names
This commit is contained in:
Mateo Wang 2026-08-12 19:48:00 -07:00 committed by GitHub
commit fdd72b5b23
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 217 additions and 14 deletions

View file

@ -1,7 +1,7 @@
# What is this?
## Helper utilities
import copy
from collections.abc import Iterable
from collections.abc import Iterable, Mapping
from typing import TYPE_CHECKING, Any, Final, Literal
import httpx
@ -181,7 +181,7 @@ def add_missing_spend_metadata_to_litellm_metadata(litellm_metadata: dict, metad
def get_metadata_variable_name_from_kwargs(
kwargs: dict,
kwargs: Mapping[str, object],
) -> Literal["metadata", "litellm_metadata"]:
"""
Helper to return what the "metadata" field should be called in the request data

View file

@ -10715,6 +10715,14 @@ class Router:
return None
@staticmethod
def _is_strategy_marker_deployment(deployment: Mapping[str, object]) -> bool:
litellm_params: Final = deployment.get("litellm_params")
if not isinstance(litellm_params, Mapping):
return False
deployment_model: Final = litellm_params.get("model")
return isinstance(deployment_model, str) and classify_strategy_router_model(deployment_model) is not None
def _common_checks_available_deployment(
self,
model: str,
@ -10842,7 +10850,12 @@ class Router:
model
] # update the model to the actual value if an alias has been passed in
return model, healthy_deployments
marker_flags: Final = tuple(self._is_strategy_marker_deployment(d) for d in healthy_deployments)
if all(marker_flags) or not any(marker_flags):
return model, healthy_deployments
return model, [ # mutable-ok: matches this function's list contract expected by downstream filters
d for d, is_marker in zip(healthy_deployments, marker_flags, strict=True) if not is_marker
]
def _filter_deployments_by_model_access_groups(
self,
@ -11355,6 +11368,10 @@ class Router:
return filtered
def _model_name_has_plain_deployments(self, model: str) -> bool:
indices: Final = self.model_name_to_deployment_indices.get(model) or ()
return any(not self._is_strategy_marker_deployment(self.model_list[idx]) for idx in indices)
def _select_pre_routing_strategy(
self, model: str, request_kwargs: dict
) -> "TaggedPreRoutingStrategy[PreRoutingStrategy] | None":
@ -11366,6 +11383,11 @@ class Router:
request's tags were what selected it, and can locate the marker
deployment the strategy was registered from via its (model_name, tags)
pair.
With tag filtering enabled, strategies that all carry real tags matching
none of the request's do not capture it when the name also has plain
deployments: returning None hands the request to ordinary tag-aware
deployment selection.
"""
candidates: Final[list[TaggedPreRoutingStrategy[PreRoutingStrategy]]] = [
*self.auto_routers.get(model, []),
@ -11375,8 +11397,6 @@ class Router:
]
if not candidates:
return None
if len(candidates) == 1:
return candidates[0]
request_tags: Final = _get_tags_from_request_kwargs(request_kwargs)
if request_tags:
@ -11388,6 +11408,12 @@ class Router:
for tagged in candidates:
if "default" in tagged.tags:
return tagged
if (
self.enable_tag_filtering
and all(tagged.tags for tagged in candidates)
and self._model_name_has_plain_deployments(model)
):
return None
return candidates[0]
async def async_pre_routing_hook(

View file

@ -584,8 +584,26 @@ async def get_deployments_for_tag(
return healthy_deployments
def _tags_in_metadata(metadata: object) -> list[str]:
"""
Tags out of a metadata bucket the caller controls the shape of.
A request can send its metadata (and its ``tags``) as anything the JSON body
allowed, an unparsed string or null included, so any shape that is not a list
of string tags carries no tags rather than raising.
"""
if not isinstance(metadata, Mapping):
return []
typed_metadata: Final[Mapping[str, object]] = metadata
tags: Final = typed_metadata.get("tags")
if isinstance(tags, str) or not isinstance(tags, Sequence):
return []
typed_tags: Final[Sequence[object]] = tags
return [tag for tag in typed_tags if isinstance(tag, str)]
def _get_tags_from_request_kwargs(
request_kwargs: dict[Any, Any] | None = None,
request_kwargs: Mapping[Any, Any] | None = None,
metadata_variable_name: Literal["metadata", "litellm_metadata"] | None = None,
) -> list[str]:
"""
@ -604,12 +622,11 @@ def _get_tags_from_request_kwargs(
return []
resolved_variable_name: Final = metadata_variable_name or get_metadata_variable_name_from_kwargs(request_kwargs)
if resolved_variable_name in request_kwargs:
metadata: Final = request_kwargs[resolved_variable_name] or {}
tags = metadata.get("tags", [])
return tags if tags is not None else []
elif "litellm_params" in request_kwargs:
litellm_params: Final = request_kwargs["litellm_params"] or {}
_metadata: Final = litellm_params.get(resolved_variable_name, {}) or {}
tags = _metadata.get("tags", [])
return tags if tags is not None else []
return _tags_in_metadata(request_kwargs[resolved_variable_name])
if "litellm_params" in request_kwargs:
litellm_params: Final = request_kwargs["litellm_params"]
if not isinstance(litellm_params, Mapping):
return []
typed_litellm_params: Final[Mapping[str, object]] = litellm_params
return _tags_in_metadata(typed_litellm_params.get(resolved_variable_name))
return []

View file

@ -1176,6 +1176,41 @@ class TestPreRoutingStrategyRegistry:
}
assert router._select_pre_routing_strategy("smart", {}).strategy is cn
@staticmethod
def _router_with_plain_smart_deployment(enable_tag_filtering: bool) -> Router:
return Router(
model_list=[{"model_name": "smart", "litellm_params": {"model": "openai/gpt-4o-mini"}}],
enable_tag_filtering=enable_tag_filtering,
)
def test_select_falls_through_to_plain_deployments_when_no_tag_matches_under_tag_filtering(self):
router = self._router_with_plain_smart_deployment(enable_tag_filtering=True)
cn, us = object(), object()
router.complexity_routers = {"smart": [TaggedPreRoutingStrategy(tags=("cn",), strategy=cn)]}
assert router._select_pre_routing_strategy("smart", {}) is None
assert router._select_pre_routing_strategy("smart", {"metadata": {"tags": ["cn"]}}).strategy is cn
router.complexity_routers = {
"smart": [
TaggedPreRoutingStrategy(tags=("cn",), strategy=cn),
TaggedPreRoutingStrategy(tags=("us",), strategy=us),
]
}
assert router._select_pre_routing_strategy("smart", {}) is None
assert router._select_pre_routing_strategy("smart", {"metadata": {"tags": ["row"]}}) is None
assert router._select_pre_routing_strategy("smart", {"metadata": {"tags": ["us"]}}).strategy is us
router.complexity_routers["router-only"] = [TaggedPreRoutingStrategy(tags=("cn",), strategy=cn)]
assert router._select_pre_routing_strategy("router-only", {}).strategy is cn
def test_select_keeps_capturing_when_tag_filtering_is_disabled(self):
router = self._router_with_plain_smart_deployment(enable_tag_filtering=False)
cn = object()
router.complexity_routers = {"smart": [TaggedPreRoutingStrategy(tags=("cn",), strategy=cn)]}
assert router._select_pre_routing_strategy("smart", {}).strategy is cn
class TestAsyncPreRoutingHookMultiFormat:
"""Test async_pre_routing_hook with multiple input formats."""

View file

@ -423,6 +423,33 @@ def test_get_tags_from_request_kwargs_various_inputs():
assert _get_tags_from_request_kwargs({"foo": "bar"}) == []
@pytest.mark.parametrize(
"request_kwargs",
[
{"metadata": "not-a-dict"},
{"litellm_metadata": "not-a-dict"},
{"litellm_metadata": ["not", "a", "dict"]},
{"litellm_params": "not-a-dict"},
{"litellm_params": {"metadata": "not-a-dict"}},
{"metadata": {"tags": "free"}},
{"metadata": {"tags": {"free": "paid"}}},
],
)
def test_get_tags_from_request_kwargs_reads_no_tags_from_a_non_dict_shape(request_kwargs):
"""Metadata and `tags` are request-controlled, so a client can send either as a
string, a list or null. Every shape that cannot hold string tags reads as untagged
instead of raising, because callers run on the hot request path."""
from litellm.router_strategy.tag_based_routing import _get_tags_from_request_kwargs
assert _get_tags_from_request_kwargs(request_kwargs) == []
def test_get_tags_from_request_kwargs_keeps_only_string_tags():
from litellm.router_strategy.tag_based_routing import _get_tags_from_request_kwargs
assert _get_tags_from_request_kwargs({"metadata": {"tags": ["free", 7, None, "paid"]}}) == ["free", "paid"]
# --- _split_tags unit tests ---

View file

@ -7617,6 +7617,104 @@ class TestAutoRouterMaxInputCharsWiring:
assert self._registered_auto_router(router).max_input_chars == DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS
class TestTaggedAutoRouterOnSharedModelName:
"""A tagged auto-router marker sharing its model_name with a plain deployment must not
capture requests whose tags don't match it when tag filtering is enabled (#36620)."""
class _FixedRouteLayer:
def __call__(self, text: str):
from semantic_router.schema import RouteChoice
return RouteChoice(name="gemini-flash")
@classmethod
def _router(cls, marker_tags, include_plain_sibling: bool, enable_tag_filtering: bool) -> "litellm.Router":
pytest.importorskip("semantic_router", reason="auto-router needs the semantic-router extra")
marker = {
"model_name": "gpt4o",
"litellm_params": {
"model": "auto_router/gpt4o-router",
"auto_router_config": json.dumps(
{"routes": [{"name": "gemini-flash", "utterances": ["capital city questions"]}]}
),
"auto_router_default_model": "gemini-flash",
"auto_router_embedding_model": "text-embedding-3-small",
**({"tags": marker_tags} if marker_tags else {}),
},
}
plain = {"model_name": "gpt4o", "litellm_params": {"model": "openai/gpt-4o"}}
tier = {"model_name": "gemini-flash", "litellm_params": {"model": "gemini/gemini-3.6-flash"}}
router = litellm.Router(
model_list=[plain, marker, tier] if include_plain_sibling else [marker, tier],
enable_tag_filtering=enable_tag_filtering,
)
router.auto_routers["gpt4o"][0].strategy.routelayer = cls._FixedRouteLayer()
return router
@staticmethod
async def _hook_response(router: "litellm.Router", request_kwargs: dict):
return await router.async_pre_routing_hook(
model="gpt4o",
request_kwargs=request_kwargs,
messages=[{"role": "user", "content": "What is the capital of France?"}],
)
@pytest.mark.asyncio
async def test_untagged_request_bypasses_the_tagged_marker_when_a_plain_deployment_shares_the_name(self):
router = self._router(marker_tags=["route"], include_plain_sibling=True, enable_tag_filtering=True)
assert await self._hook_response(router, {}) is None
@pytest.mark.asyncio
async def test_request_tagged_for_the_marker_is_still_semantically_routed(self):
router = self._router(marker_tags=["route"], include_plain_sibling=True, enable_tag_filtering=True)
response = await self._hook_response(router, {"metadata": {"tags": ["route"]}})
assert response is not None
assert response.model == "gemini-flash"
@pytest.mark.asyncio
async def test_marker_only_alias_still_captures_untagged_requests(self):
router = self._router(marker_tags=["route"], include_plain_sibling=False, enable_tag_filtering=True)
response = await self._hook_response(router, {})
assert response is not None
assert response.model == "gemini-flash"
@pytest.mark.asyncio
async def test_untagged_marker_sharing_the_name_still_captures_untagged_requests(self):
router = self._router(marker_tags=None, include_plain_sibling=True, enable_tag_filtering=True)
response = await self._hook_response(router, {})
assert response is not None
assert response.model == "gemini-flash"
@pytest.mark.asyncio
async def test_untagged_selection_never_lands_on_the_marker_deployment(self):
router = self._router(marker_tags=["route"], include_plain_sibling=True, enable_tag_filtering=True)
for _ in range(20):
deployment = await router.async_get_available_deployment(
model="gpt4o",
request_kwargs={},
messages=[{"role": "user", "content": "What is the capital of France?"}],
)
assert deployment["litellm_params"]["model"] == "openai/gpt-4o"
def test_deployment_without_litellm_params_mapping_is_not_a_marker(self):
assert litellm.Router._is_strategy_marker_deployment({"model_name": "gpt4o"}) is False
def test_model_name_has_plain_deployments_reflects_the_pool(self):
mixed = self._router(marker_tags=["route"], include_plain_sibling=True, enable_tag_filtering=True)
marker_only = self._router(marker_tags=["route"], include_plain_sibling=False, enable_tag_filtering=True)
assert mixed._model_name_has_plain_deployments("gpt4o") is True
assert marker_only._model_name_has_plain_deployments("gpt4o") is False
class TestGetAllowedFailsFromPolicy:
def _make_router(self, **policy_kwargs) -> litellm.Router:
from litellm.types.router import AllowedFailsPolicy