mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
fix(router): support tag-based routing for complexity_router deployments sharing a model_name
This commit is contained in:
parent
4d33964898
commit
73524df72a
3 changed files with 160 additions and 10 deletions
|
|
@ -488,7 +488,7 @@ class Router:
|
|||
self.pattern_router = PatternMatchRouter()
|
||||
self.team_pattern_routers: Dict[str, PatternMatchRouter] = {} # {"TEAM_ID": PatternMatchRouter}
|
||||
self.auto_routers: Dict[str, "AutoRouter"] = {}
|
||||
self.complexity_routers: Dict[str, "ComplexityRouter"] = {}
|
||||
self.complexity_routers: Dict[str, list["ComplexityRouter"]] = {}
|
||||
self.adaptive_routers: Dict[str, "AdaptiveRouter"] = {}
|
||||
self.quality_routers: Dict[str, "QualityRouter"] = {}
|
||||
self.routing_plugins: list[RoutingPlugin] = list(plugins) if plugins else []
|
||||
|
|
@ -7652,17 +7652,22 @@ class Router:
|
|||
"or configure tiers in complexity_router_config. Please set it in the litellm_params"
|
||||
)
|
||||
|
||||
deployment_tags = deployment.litellm_params.tags or []
|
||||
complexity_router: ComplexityRouter = ComplexityRouter(
|
||||
model_name=deployment.model_name,
|
||||
default_model=default_model,
|
||||
litellm_router_instance=self,
|
||||
complexity_router_config=complexity_router_config,
|
||||
deployment_tags=deployment_tags,
|
||||
)
|
||||
if deployment.model_name in self.complexity_routers:
|
||||
existing_routers = self.complexity_routers.get(deployment.model_name, [])
|
||||
new_tag_set = frozenset(deployment_tags)
|
||||
if any(frozenset(router.deployment_tags) == new_tag_set for router in existing_routers):
|
||||
raise ValueError(
|
||||
f"Complexity-router deployment {deployment.model_name} already exists. Please use a different model name."
|
||||
f"Complexity-router deployment {deployment.model_name} already exists with tags "
|
||||
f"{sorted(new_tag_set)}. Give each deployment sharing a model_name a distinct set of tags."
|
||||
)
|
||||
self.complexity_routers[deployment.model_name] = complexity_router
|
||||
self.complexity_routers[deployment.model_name] = [*existing_routers, complexity_router]
|
||||
|
||||
def _is_adaptive_router_deployment(self, litellm_params: LiteLLM_Params) -> bool:
|
||||
"""True when this deployment opts in via the `auto_router/adaptive_router` model prefix."""
|
||||
|
|
@ -7698,12 +7703,16 @@ class Router:
|
|||
continue
|
||||
self.init_adaptive_router_deployment(deployment=deployment)
|
||||
|
||||
for model_name, complexity_router in self.complexity_routers.items():
|
||||
if not complexity_router.config.adaptive or model_name in self.adaptive_routers:
|
||||
for model_name, complexity_router_list in self.complexity_routers.items():
|
||||
if model_name in self.adaptive_routers:
|
||||
continue
|
||||
adaptive_router = complexity_router._ensure_adaptive_router()
|
||||
if adaptive_router is not None:
|
||||
self.adaptive_routers[model_name] = adaptive_router
|
||||
for complexity_router in complexity_router_list:
|
||||
if not complexity_router.config.adaptive:
|
||||
continue
|
||||
adaptive_router = complexity_router._ensure_adaptive_router()
|
||||
if adaptive_router is not None:
|
||||
self.adaptive_routers[model_name] = adaptive_router
|
||||
break
|
||||
|
||||
for callback in litellm.logging_callback_manager.get_custom_loggers_for_type(AdaptiveRouterPostCallHook):
|
||||
litellm.logging_callback_manager.remove_callback_from_all_lists(callback)
|
||||
|
|
@ -10810,6 +10819,44 @@ class Router:
|
|||
|
||||
return filtered
|
||||
|
||||
def _select_complexity_router(self, model: str, request_kwargs: dict) -> Optional["ComplexityRouter"]:
|
||||
"""
|
||||
Resolve which complexity router handles this request.
|
||||
|
||||
Multiple complexity-router deployments can share a model_name while
|
||||
targeting different tag groups (e.g. one per region). Pick the router
|
||||
whose deployment tags match the request's tags, falling back to a
|
||||
`default`-tagged router and finally to the first registered one so
|
||||
single-deployment setups keep their existing behaviour.
|
||||
"""
|
||||
from litellm.router_strategy.tag_based_routing import (
|
||||
_get_tags_from_request_kwargs,
|
||||
is_valid_deployment_tag,
|
||||
)
|
||||
|
||||
routers = self.complexity_routers.get(model)
|
||||
if not routers:
|
||||
return None
|
||||
if len(routers) == 1:
|
||||
return routers[0]
|
||||
|
||||
request_tags = _get_tags_from_request_kwargs(
|
||||
request_kwargs=request_kwargs,
|
||||
metadata_variable_name=self._get_metadata_variable_name_from_kwargs(request_kwargs),
|
||||
)
|
||||
if request_tags:
|
||||
for router in routers:
|
||||
if router.deployment_tags and is_valid_deployment_tag(
|
||||
router.deployment_tags, request_tags, self.tag_filtering_match_any
|
||||
):
|
||||
return router
|
||||
|
||||
for router in routers:
|
||||
if "default" in router.deployment_tags:
|
||||
return router
|
||||
|
||||
return routers[0]
|
||||
|
||||
async def async_pre_routing_hook(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -10834,7 +10881,7 @@ class Router:
|
|||
|
||||
router_strategy = (
|
||||
self.auto_routers.get(model)
|
||||
or self.complexity_routers.get(model)
|
||||
or self._select_complexity_router(model=model, request_kwargs=request_kwargs)
|
||||
or self.adaptive_routers.get(model)
|
||||
or self.quality_routers.get(model)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -142,6 +142,7 @@ class ComplexityRouter(CustomLogger):
|
|||
litellm_router_instance: Router,
|
||||
complexity_router_config: dict[str, Any] | None = None,
|
||||
default_model: str | None = None,
|
||||
deployment_tags: list[str] | None = None,
|
||||
):
|
||||
"""
|
||||
Initialize ComplexityRouter.
|
||||
|
|
@ -151,9 +152,13 @@ class ComplexityRouter(CustomLogger):
|
|||
litellm_router_instance: The LiteLLM Router instance.
|
||||
complexity_router_config: Optional configuration dict from proxy config.
|
||||
default_model: Optional default model to use if tier cannot be determined.
|
||||
deployment_tags: Tags configured on the deployment this router backs.
|
||||
Used to disambiguate between multiple complexity-router deployments
|
||||
that share a model_name but target different tag groups.
|
||||
"""
|
||||
self.model_name = model_name
|
||||
self.litellm_router_instance = litellm_router_instance
|
||||
self.deployment_tags = deployment_tags or []
|
||||
|
||||
# Parse config - always create a new instance to avoid singleton mutation
|
||||
if complexity_router_config:
|
||||
|
|
|
|||
|
|
@ -2984,3 +2984,101 @@ class TestRoutingPlugins:
|
|||
assert first.model == "gpt-4o-mini"
|
||||
assert second.model == "gpt-4o-mini"
|
||||
assert spy.call_count == 2
|
||||
|
||||
|
||||
class TestComplexityRouterTagBasedRouting:
|
||||
"""Regression tests for https://github.com/BerriAI/litellm/issues/33655.
|
||||
|
||||
Two complexity-router deployments that share a model_name but carry different
|
||||
tags must both be registered, and each request must resolve to the config
|
||||
whose tags match the request's tags.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def _two_region_model_list():
|
||||
return [
|
||||
{
|
||||
"model_name": "smart-router",
|
||||
"litellm_params": {
|
||||
"model": "auto_router/complexity_router",
|
||||
"complexity_router_default_model": "cn/simple",
|
||||
"complexity_router_config": {
|
||||
"tiers": {
|
||||
"SIMPLE": "cn/simple",
|
||||
"MEDIUM": "cn/simple",
|
||||
"COMPLEX": "cn/complex",
|
||||
"REASONING": "cn/complex",
|
||||
}
|
||||
},
|
||||
"tags": ["cn-dev"],
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "smart-router",
|
||||
"litellm_params": {
|
||||
"model": "auto_router/complexity_router",
|
||||
"complexity_router_default_model": "row/simple",
|
||||
"complexity_router_config": {
|
||||
"tiers": {
|
||||
"SIMPLE": "row/simple",
|
||||
"MEDIUM": "row/simple",
|
||||
"COMPLEX": "row/complex",
|
||||
"REASONING": "row/complex",
|
||||
}
|
||||
},
|
||||
"tags": ["row-dev"],
|
||||
},
|
||||
},
|
||||
]
|
||||
|
||||
def test_both_tag_differentiated_deployments_registered(self):
|
||||
router = Router(model_list=self._two_region_model_list())
|
||||
|
||||
registered = router.complexity_routers["smart-router"]
|
||||
assert len(registered) == 2
|
||||
assert [r.deployment_tags for r in registered] == [["cn-dev"], ["row-dev"]]
|
||||
|
||||
def test_duplicate_tag_set_still_rejected(self):
|
||||
model_list = self._two_region_model_list()
|
||||
model_list[1]["litellm_params"]["tags"] = ["cn-dev"]
|
||||
|
||||
with pytest.raises(ValueError, match="already exists with tags"):
|
||||
Router(model_list=model_list)
|
||||
|
||||
def test_select_router_by_request_tags(self):
|
||||
router = Router(model_list=self._two_region_model_list())
|
||||
|
||||
cn = router._select_complexity_router(
|
||||
model="smart-router", request_kwargs={"metadata": {"tags": ["cn-dev"]}}
|
||||
)
|
||||
row = router._select_complexity_router(
|
||||
model="smart-router", request_kwargs={"metadata": {"tags": ["row-dev"]}}
|
||||
)
|
||||
|
||||
assert cn is not None and cn.deployment_tags == ["cn-dev"]
|
||||
assert row is not None and row.deployment_tags == ["row-dev"]
|
||||
|
||||
def test_select_router_falls_back_to_first_without_tags(self):
|
||||
router = Router(model_list=self._two_region_model_list())
|
||||
|
||||
selected = router._select_complexity_router(model="smart-router", request_kwargs={})
|
||||
|
||||
assert selected is not None and selected.deployment_tags == ["cn-dev"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_routing_hook_routes_to_matching_region_tier(self):
|
||||
router = Router(model_list=self._two_region_model_list())
|
||||
|
||||
cn_response = await router.async_pre_routing_hook(
|
||||
model="smart-router",
|
||||
request_kwargs={"metadata": {"tags": ["cn-dev"]}},
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
)
|
||||
row_response = await router.async_pre_routing_hook(
|
||||
model="smart-router",
|
||||
request_kwargs={"metadata": {"tags": ["row-dev"]}},
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
)
|
||||
|
||||
assert cn_response is not None and cn_response.model == "cn/simple"
|
||||
assert row_response is not None and row_response.model == "row/simple"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue