From c5dea3406432dafdfb02a0713e55c7e54fe3393a Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Thu, 12 Dec 2024 15:19:50 -0800 Subject: [PATCH] fix(router.py): fix pattern matching router - add generic "*" to it as well Fixes issue where generic "*" model access group wouldn't show up --- litellm/proxy/_new_secret_config.yaml | 26 +++-------------- litellm/proxy/auth/auth_checks.py | 2 +- litellm/proxy/proxy_server.py | 9 +++--- litellm/router.py | 14 +++++----- .../router_utils/pattern_match_deployments.py | 1 - .../test_router_pattern_matching.py | 28 +++++++++++++++++++ 6 files changed, 44 insertions(+), 36 deletions(-) diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index 84075f53e05..a66057ae301 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -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 . My number is: " - items: - - start: 48 - end: 62 - entity_type: PHONE_NUMBER - text: "" - operator: replace - - start: 24 - end: 32 - entity_type: PERSON - text: "" - operator: replace - -litellm_settings: - set_verbose: true - success_callback: ["langfuse"] \ No newline at end of file + model: "*" + model_info: + access_groups: ["default"] \ No newline at end of file diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 3b43ec32e11..b74e5199e8e 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -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 diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 93df33d7574..cb540c5f05e 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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") diff --git a/litellm/router.py b/litellm/router.py index 2f333bf6b38..85bd63bf2a4 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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( diff --git a/litellm/router_utils/pattern_match_deployments.py b/litellm/router_utils/pattern_match_deployments.py index a369100eb09..3b10dcdb769 100644 --- a/litellm/router_utils/pattern_match_deployments.py +++ b/litellm/router_utils/pattern_match_deployments.py @@ -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 diff --git a/tests/local_testing/test_router_pattern_matching.py b/tests/local_testing/test_router_pattern_matching.py index 914e8ecfa9d..35e4d2c3d84 100644 --- a/tests/local_testing/test_router_pattern_matching.py +++ b/tests/local_testing/test_router_pattern_matching.py @@ -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?"}], + )