mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
27544e4328
commit
c5dea34064
6 changed files with 44 additions and 36 deletions
|
|
@ -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"]
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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?"}],
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue