This commit is contained in:
Souravrajvi0 2026-08-27 16:33:57 -04:00 • committed by GitHub
commit 46e148c9d7
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 131 additions and 0 deletions

View file

@ -179,6 +179,7 @@ from litellm.router_utils.router_callbacks.track_deployment_metrics import (
increment_deployment_successes_for_current_minute,
)
from litellm.scheduler import FlowItem, Scheduler
from litellm.secret_managers.main import get_secret_bool
from litellm.types.llms.openai import (
AllMessageValues,
FileTypes,
@ -604,6 +605,7 @@ class Router:
health_check_ignore_transient_errors: bool = False,
background_health_check_model_groups: Sequence[str] | None = None,
enable_weighted_failover: bool = False,
enable_cross_model_group_collision_check: bool = False,
) -> None:
"""
Initialize the Router class with the given parameters for caching, reliability, and routing strategy.
@ -640,6 +642,7 @@ class Router:
deployment_affinity_ttl_seconds (int): TTL for user-key -> deployment affinity mapping. Defaults to 3600.
ignore_invalid_deployments (bool): Ignores invalid deployments, and continues with other deployments. Default is to raise an error.
enable_weighted_failover (bool): When True and the routing strategy is "simple-shuffle", a retryable failure on one deployment causes the request to re-pick (weighted) across the other deployments in the same model group before any cross-group fallback runs. Bounded by `max_fallbacks`. Async-only: currently honored by `router.acompletion()` and other async entrypoints. The sync `router.completion()` path falls back to the regular fallback flow. Defaults to False.
enable_cross_model_group_collision_check (bool): When True, refuse the fallback that load-balances a raw `litellm_params.model` string across more than one `model_name` group. Opt-in; also enabled by `LITELLM_ENABLE_CROSS_MODEL_GROUP_COLLISION_CHECK`. Defaults to False.
Returns:
Router: An instance of the litellm.Router class.
@ -811,6 +814,10 @@ class Router:
self.disable_cooldowns = disable_cooldowns
self.enable_health_check_routing = enable_health_check_routing
self.enable_weighted_failover = enable_weighted_failover
self.enable_cross_model_group_collision_check = bool(
enable_cross_model_group_collision_check
or get_secret_bool("LITELLM_ENABLE_CROSS_MODEL_GROUP_COLLISION_CHECK", False)
)
self.health_check_ignore_transient_errors = health_check_ignore_transient_errors
self.background_health_check_model_groups: frozenset[str] | None = (
frozenset(background_health_check_model_groups)
@ -10909,6 +10916,7 @@ class Router:
"retry_policy",
"model_group_alias",
"enable_weighted_failover",
"enable_cross_model_group_collision_check",
"enable_tag_filtering",
"tag_routing_prefix",
]
@ -10947,6 +10955,7 @@ class Router:
"model_group_retry_policy",
"model_group_alias",
"enable_weighted_failover",
"enable_cross_model_group_collision_check",
"enable_tag_filtering",
"tag_routing_prefix",
]
@ -11341,6 +11350,24 @@ class Router:
"""
return [m for m in self.model_list if m["litellm_params"]["model"] == model]
def _reject_cross_model_group_collision(self, model: str, deployments: list) -> None:
group_names: Final = frozenset(
str(deployment["model_name"])
for deployment in deployments
if isinstance(deployment.get("model_name"), str) and deployment["model_name"]
)
if len(group_names) <= 1:
return
groups: Final = ", ".join(sorted(group_names))
raise litellm.BadRequestError(
message=(
f"Model '{model}' matches deployments from multiple model groups ({groups}). "
"Request a configured model_name instead of the raw provider model string."
),
model=model,
llm_provider="",
)
def _try_early_resolve_deployments_for_model_not_in_names(
self,
model: str,
@ -11496,6 +11523,8 @@ class Router:
request_kwargs=request_kwargs,
request_team_id=request_team_id,
)
if self.enable_cross_model_group_collision_check:
self._reject_cross_model_group_collision(model=model, deployments=healthy_deployments)
# If the litellm-model lookup produced candidates that access-group
# filtering then removed, treat this the same as the by-name path
# being emptied: prevent default-model fallback from bypassing the

View file

@ -220,6 +220,14 @@ ROUTER_SETTINGS_FIELDS: Final[list[RouterSettingsField]] = [
field_default=False,
ui_field_name="Enable Pre-call Checks",
),
RouterSettingsField(
field_name="enable_cross_model_group_collision_check",
field_type="Boolean",
field_value=None,
field_description=("Reject raw provider model strings that match more than one model_name group"),
field_default=False,
ui_field_name="Cross Model Group Collision Check",
),
RouterSettingsField(
field_name="default_litellm_params",
field_type="Dictionary",

View file

@ -125,6 +125,7 @@ class UpdateRouterConfig(BaseModel):
model_group_alias: dict[str, str | dict] | None = {}
enable_tag_filtering: bool | None = None
tag_routing_prefix: str | None = None
enable_cross_model_group_collision_check: bool | None = None
model_config = ConfigDict(protected_namespaces=())

View file

@ -10669,3 +10669,96 @@ def test_permission_denied_error_is_retried_when_other_deployments_exist():
)
is True
)
def _colliding_raw_model_groups() -> list[dict]:
shared_model = "openai/some-shared-string"
return [
{
"model_name": "chain-a",
"litellm_params": {
"model": shared_model,
"api_key": "key-a",
"mock_response": "from-a",
},
},
{
"model_name": "chain-b",
"litellm_params": {
"model": shared_model,
"api_key": "key-b",
"mock_response": "from-b",
},
},
]
def test_raw_litellm_model_fallback_spans_unrelated_groups_by_default():
"""Issue #38216: a raw litellm_params.model string load-balances across every group that shares it."""
router = litellm.Router(model_list=_colliding_raw_model_groups())
_model, deployments = router._common_checks_available_deployment(
model="openai/some-shared-string"
)
assert {d["model_name"] for d in deployments} == {"chain-a", "chain-b"}
assert router.get_available_deployment(model="openai/some-shared-string") is not None
def test_cross_model_group_collision_check_rejects_ambiguous_raw_model():
router = litellm.Router(
model_list=_colliding_raw_model_groups(),
enable_cross_model_group_collision_check=True,
)
with pytest.raises(litellm.BadRequestError, match="multiple model groups"):
router.get_available_deployment(model="openai/some-shared-string")
deployment = router.get_available_deployment(model="chain-a")
assert deployment["model_name"] == "chain-a"
def test_cross_model_group_collision_check_allows_same_group_and_unambiguous_raw_model():
router = litellm.Router(
model_list=[
{
"model_name": "chain-a",
"litellm_params": {
"model": "openai/some-shared-string",
"api_key": "key-a",
"mock_response": "from-a-1",
},
},
{
"model_name": "chain-a",
"litellm_params": {
"model": "openai/some-shared-string",
"api_key": "key-a-2",
"mock_response": "from-a-2",
},
},
{
"model_name": "chain-b",
"litellm_params": {
"model": "openai/other-string",
"api_key": "key-b",
"mock_response": "from-b",
},
},
],
enable_cross_model_group_collision_check=True,
)
_model, deployments = router._common_checks_available_deployment(
model="openai/some-shared-string"
)
assert {d["model_name"] for d in deployments} == {"chain-a"}
assert len(deployments) == 2
def test_cross_model_group_collision_check_env_enables_guard(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("LITELLM_ENABLE_CROSS_MODEL_GROUP_COLLISION_CHECK", "true")
router = litellm.Router(model_list=_colliding_raw_model_groups())
with pytest.raises(litellm.BadRequestError, match="multiple model groups"):
router.get_available_deployment(model="openai/some-shared-string")