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
This commit is contained in:
Krrish Dholakia 2024-12-12 15:36:23 -08:00
parent c5dea34064
commit 2f9fbaf421
2 changed files with 68 additions and 3 deletions

View file

@ -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

View file

@ -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/*"