From 2f9fbaf421d462bc75f6df0b55c2e009618ce2f3 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Thu, 12 Dec 2024 15:36:23 -0800 Subject: [PATCH] fix(pattern_match_deployments.py): match to more specific pattern match to more specific pattern allows setting generic wildcard model access group and excluding specific models more easily --- .../router_utils/pattern_match_deployments.py | 43 ++++++++++++++++++- .../test_router_pattern_matching.py | 28 +++++++++++- 2 files changed, 68 insertions(+), 3 deletions(-) diff --git a/litellm/router_utils/pattern_match_deployments.py b/litellm/router_utils/pattern_match_deployments.py index 3b10dcdb769..63e8a0c42ea 100644 --- a/litellm/router_utils/pattern_match_deployments.py +++ b/litellm/router_utils/pattern_match_deployments.py @@ -4,13 +4,51 @@ Class to handle llm wildcard routing and regex pattern matching import copy import re +from functools import cached_property from re import Match -from typing import Dict, List, Optional +from typing import Dict, List, Optional, Tuple from litellm import get_llm_provider from litellm._logging import verbose_router_logger +class PatternUtils: + @staticmethod + def calculate_pattern_specificity(pattern: str) -> Tuple[int, int]: + """ + Calculate pattern specificity based on length and complexity. + + Args: + pattern: Regex pattern to analyze + + Returns: + Tuple of (length, complexity) for sorting + """ + complexity_chars = ["*", "+", "?", "\\", "^", "$", "|", "(", ")"] + return ( + len(pattern), # Longer patterns more specific + sum( + pattern.count(char) for char in complexity_chars + ), # More regex complexity + ) + + @staticmethod + def sorted_patterns( + patterns: Dict[str, List[Dict]] + ) -> List[Tuple[str, List[Dict]]]: + """ + Cached property for patterns sorted by specificity. + + Returns: + Sorted list of pattern-deployment tuples + """ + return sorted( + patterns.items(), + key=lambda x: PatternUtils.calculate_pattern_specificity(x[0]), + reverse=True, + ) + + class PatternMatchRouter: """ Class to handle llm wildcard routing and regex pattern matching @@ -99,12 +137,13 @@ class PatternMatchRouter: if request is None: return None + sorted_patterns = PatternUtils.sorted_patterns(self.patterns) regex_filtered_model_names = ( [self._pattern_to_regex(m) for m in filtered_model_names] if filtered_model_names is not None else [] ) - for pattern, llm_deployments in self.patterns.items(): + for pattern, llm_deployments in sorted_patterns: if ( filtered_model_names is not None and pattern not in regex_filtered_model_names diff --git a/tests/local_testing/test_router_pattern_matching.py b/tests/local_testing/test_router_pattern_matching.py index 35e4d2c3d84..619c594421c 100644 --- a/tests/local_testing/test_router_pattern_matching.py +++ b/tests/local_testing/test_router_pattern_matching.py @@ -133,7 +133,7 @@ def test_route_with_multiple_matching_patterns(): router.add_pattern("openai/*", deployment1.to_json(exclude_none=True)) router.add_pattern("openai/gpt-*", deployment2.to_json(exclude_none=True)) assert router.route("openai/gpt-3.5-turbo") == [ - deployment1.to_json(exclude_none=True) + deployment2.to_json(exclude_none=True) ] @@ -265,3 +265,29 @@ def test_pattern_matching_router_with_default_wildcard(): model="gpt-3.5-turbo", messages=[{"role": "user", "content": "Hello, how are you?"}], ) + + +def test_pattern_matching_router_with_default_wildcard_and_model_wildcard(): + """ + Match to more specific pattern first. + """ + router = Router( + model_list=[ + { + "model_name": "*", + "litellm_params": {"model": "*"}, + "model_info": {"access_groups": ["default"]}, + }, + { + "model_name": "llmengine/*", + "litellm_params": {"model": "openai/*"}, + }, + ] + ) + + assert len(router.pattern_router.patterns) > 0 + + pattern_router = router.pattern_router + deployments = pattern_router.route("llmengine/gpt-3.5-turbo") + assert len(deployments) == 1 + assert deployments[0]["model_name"] == "llmengine/*"