fix(router.py): fix pattern matching router - add generic "*" to it as well

Fixes issue where generic "*" model access group wouldn't show up
This commit is contained in:
Krrish Dholakia 2024-12-12 15:19:50 -08:00
parent 27544e4328
commit c5dea34064
6 changed files with 44 additions and 36 deletions

View file

@ -5,26 +5,8 @@ model_list:
api_key: os.environ/AZURE_API_KEY
api_base: os.environ/AZURE_API_BASE
temperature: 0.2
guardrails:
- guardrail_name: "presidio-log-guard"
- model_name: "*"
litellm_params:
guardrail: presidio
mode: "logging_only"
mock_redacted_text:
text: "hello world, my name is <PERSON>. My number is: <PHONE_NUMBER>"
items:
- start: 48
end: 62
entity_type: PHONE_NUMBER
text: "<PHONE_NUMBER>"
operator: replace
- start: 24
end: 32
entity_type: PERSON
text: "<PERSON>"
operator: replace
litellm_settings:
set_verbose: true
success_callback: ["langfuse"]
model: "*"
model_info:
access_groups: ["default"]

View file

@ -757,6 +757,7 @@ async def get_key_object(
except DB_CONNECTION_ERROR_TYPES as e:
return await _handle_failed_db_connection_for_get_key_object(e=e)
except Exception:
traceback.print_exc()
raise Exception(
f"Key doesn't exist in db. key={hashed_token}. Create key via `/key/generate` call."
)
@ -870,7 +871,6 @@ async def can_key_call_model(
access_groups = defaultdict(list)
if llm_router:
access_groups = llm_router.get_model_access_groups(model_name=model)
if (
len(access_groups) > 0 and llm_router is not None
): # check if token contains any model access groups

View file

@ -25,8 +25,6 @@ from typing import (
get_type_hints,
)
import requests
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
@ -120,7 +118,7 @@ from litellm.litellm_core_utils.core_helpers import (
_get_parent_otel_span_from_kwargs,
get_litellm_metadata_from_kwargs,
)
from litellm.llms.custom_httpx.httpx_handler import HTTPHandler
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.proxy._types import *
from litellm.proxy.analytics_endpoints.analytics_endpoints import (
router as analytics_router,
@ -526,7 +524,7 @@ async_result = None
celery_app_conn = None
celery_fn = None # Redis Queue for handling requests
### DB WRITER ###
db_writer_client: Optional[HTTPHandler] = None
db_writer_client: Optional[AsyncHTTPHandler] = None
### logger ###
@ -8060,7 +8058,8 @@ def get_image():
# Check if the logo path is an HTTP/HTTPS URL
if logo_path.startswith(("http://", "https://")):
# Download the image and cache it
response = requests.get(logo_path)
client = HTTPHandler()
response = client.get(logo_path)
if response.status_code == 200:
# Save the image to a local file
cache_path = os.path.join(current_dir, "cached_logo.jpg")

View file

@ -4044,15 +4044,15 @@ class Router:
# Check if user is trying to use model_name == "*"
# this is a catch all model for their specific api key
if deployment.model_name == "*":
if deployment.litellm_params.model == "*":
# user wants to pass through all requests to litellm.acompletion for unknown deployments
self.router_general_settings.pass_through_all_models = True
else:
self.default_deployment = deployment.to_json(exclude_none=True)
# if deployment.model_name == "*":
# if deployment.litellm_params.model == "*":
# # user wants to pass through all requests to litellm.acompletion for unknown deployments
# self.router_general_settings.pass_through_all_models = True
# else:
# self.default_deployment = deployment.to_json(exclude_none=True)
# Check if user is using provider specific wildcard routing
# example model_name = "databricks/*" or model_name = "anthropic/*"
elif "*" in deployment.model_name:
if "*" in deployment.model_name:
# store this as a regex pattern - all deployments matching this pattern will be sent to this deployment
# Store deployment.model_name as a regex pattern
self.pattern_router.add_pattern(

View file

@ -104,7 +104,6 @@ class PatternMatchRouter:
if filtered_model_names is not None
else []
)
for pattern, llm_deployments in self.patterns.items():
if (
filtered_model_names is not None

View file

@ -237,3 +237,31 @@ def test_router_pattern_match_e2e():
"model": "gpt-4o",
"messages": [{"role": "user", "content": "Hello, how are you?"}],
}
def test_pattern_matching_router_with_default_wildcard():
"""
Tests that the router returns the default wildcard model when the pattern is not found
Make sure generic '*' allows all models to be passed through.
"""
router = Router(
model_list=[
{
"model_name": "*",
"litellm_params": {"model": "*"},
"model_info": {"access_groups": ["default"]},
},
{
"model_name": "anthropic-claude",
"litellm_params": {"model": "anthropic/claude-3-5-sonnet"},
},
]
)
assert len(router.pattern_router.patterns) > 0
router.completion(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": "Hello, how are you?"}],
)