From fab24fae1a63cf106421f189404b070af55b06b7 Mon Sep 17 00:00:00 2001 From: Thiago Salvatore Date: Thu, 12 Jun 2025 17:18:42 -0300 Subject: [PATCH] fix: Do not add default model on tag based-routing when valid tag (#11454) * Do not add default when valid tagged model * Use default models when no tag matches * Add unit tests --- litellm/router_strategy/tag_based_routing.py | 15 ++++----- .../local_testing/test_router_tag_routing.py | 31 +++++++++++++++++-- 2 files changed, 35 insertions(+), 11 deletions(-) diff --git a/litellm/router_strategy/tag_based_routing.py b/litellm/router_strategy/tag_based_routing.py index e6a936140d7..34261d83dcf 100644 --- a/litellm/router_strategy/tag_based_routing.py +++ b/litellm/router_strategy/tag_based_routing.py @@ -33,13 +33,6 @@ def is_valid_deployment_tag( request_tags, ) return True - elif "default" in deployment_tags: - verbose_logger.debug( - "adding default deployment with tags: %s, request tags: %s", - deployment_tags, - request_tags, - ) - return True return False @@ -76,6 +69,7 @@ async def get_deployments_for_tag( request_tags = metadata.get("tags") new_healthy_deployments = [] + default_deployments = [] if request_tags: verbose_logger.debug( "get_deployments_for_tag routing: router_keys: %s", request_tags @@ -98,12 +92,15 @@ async def get_deployments_for_tag( if is_valid_deployment_tag(deployment_tags, request_tags): new_healthy_deployments.append(deployment) - if len(new_healthy_deployments) == 0: + if "default" in deployment_tags: + default_deployments.append(deployment) + + if len(new_healthy_deployments) == 0 and len(default_deployments) == 0: raise ValueError( f"{RouterErrors.no_deployments_with_tag_routing.value}. Passed model={model} and tags={request_tags}" ) - return new_healthy_deployments + return new_healthy_deployments if len(new_healthy_deployments) > 0 else default_deployments # for Untagged requests use default deployments if set _default_deployments_with_tags = [] diff --git a/tests/local_testing/test_router_tag_routing.py b/tests/local_testing/test_router_tag_routing.py index 4e30e1d8b6c..87cf2261a67 100644 --- a/tests/local_testing/test_router_tag_routing.py +++ b/tests/local_testing/test_router_tag_routing.py @@ -116,11 +116,21 @@ async def test_router_free_paid_tier_embeddings(): }, "model_info": {"id": "very-expensive-model"}, }, + { + "model_name": "gpt-4", + "litellm_params": { + "model": "gpt-4o-mini", + "api_base": "https://exampleopenaiendpoint-production.up.railway.app/", + "tags": ["default"], + "mock_response": ["1", "2", "3"], + }, + "model_info": {"id": "default-model"}, + }, ], enable_tag_filtering=True, ) - for _ in range(1): + for _ in range(5): # this should pick model with id == very-cheap-model response = await router.aembedding( model="gpt-4", @@ -136,7 +146,7 @@ async def test_router_free_paid_tier_embeddings(): assert response_extra_info["model_id"] == "very-cheap-model" for _ in range(5): - # this should pick model with id == very-cheap-model + # this should pick model with id == very-expensive-model response = await router.aembedding( model="gpt-4", input="Tell me a joke.", @@ -219,6 +229,22 @@ async def test_default_tagged_deployments(): assert response_extra_info["model_id"] == "default-model" + for _ in range(5): + # requests with invalid tags, this should pick model with id == "default-model" + response = await router.acompletion( + model="gpt-4", + messages=[{"role": "user", "content": "Tell me a joke."}], + metadata={"tags": ["invalid-tag"]}, + ) + + print("Response: ", response) + + response_extra_info = response._hidden_params + print("response_extra_info: ", response_extra_info) + + assert response_extra_info["model_id"] == "default-model" + + @pytest.mark.asyncio() async def test_error_from_tag_routing(): @@ -288,3 +314,4 @@ def test_tag_routing_with_list_of_tags(): assert is_valid_deployment_tag(["teamA", "teamB"], ["teamA", "teamC"]) assert not is_valid_deployment_tag(["teamA", "teamB"], ["teamC"]) assert not is_valid_deployment_tag(["teamA", "teamB"], []) + assert not is_valid_deployment_tag(["default"], ["teamA"])