From 962c33ee721b48beca92ee9a80a2652f3dbbbe1f Mon Sep 17 00:00:00 2001 From: Krish Dholakia Date: Wed, 9 Jul 2025 11:42:11 -0700 Subject: [PATCH] (Router) don't add invalid deployment to router pattern match (#12459) * fix security - mcp * fix(router.py): validate model provider before adding deployment to pattern prevents routing on pattern match to invalid deployment --------- Co-authored-by: Ishaan Jaff --- litellm/router.py | 41 +++++++++++++------------------ tests/test_litellm/test_router.py | 29 ++++++++++++++++++++++ 2 files changed, 46 insertions(+), 24 deletions(-) diff --git a/litellm/router.py b/litellm/router.py index 4ed86580588..5b64f8b1ee5 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -1735,7 +1735,6 @@ class Router: kwargs: Dict[str, Any], ): litellm_logging_object = kwargs.get("litellm_logging_obj", None) - if litellm_logging_object is None: litellm_logging_object, kwargs = function_setup( **{ @@ -1771,10 +1770,6 @@ class Router: "litellm_params" ].get("prompt_label", None) - prompt_version = kwargs.get( - "prompt_version", None - ) or prompt_management_deployment["litellm_params"].get("prompt_version", None) - if prompt_id is None or not isinstance(prompt_id, str): raise ValueError( f"Prompt ID is not set or not a string. Got={prompt_id}, type={type(prompt_id)}" @@ -1795,7 +1790,6 @@ class Router: prompt_id=prompt_id, prompt_variables=prompt_variables, prompt_label=prompt_label, - prompt_version=prompt_version, ) kwargs = {**data, **kwargs, **optional_params} @@ -4723,6 +4717,23 @@ class Router: def _add_deployment(self, deployment: Deployment) -> Deployment: import os + #### VALIDATE MODEL ######## + # check if model provider in supported providers + ( + _model, + custom_llm_provider, + dynamic_api_key, + api_base, + ) = litellm.get_llm_provider( + model=deployment.litellm_params.model, + custom_llm_provider=deployment.litellm_params.get( + "custom_llm_provider", None + ), + ) + # done reading model["litellm_params"] + if custom_llm_provider not in litellm.provider_list: + raise Exception(f"Unsupported provider - {custom_llm_provider}") + #### DEPLOYMENT NAMES INIT ######## self.deployment_names.append(deployment.litellm_params.model) ############ Users can either pass tpm/rpm as a litellm_param or a router param ########### @@ -4740,20 +4751,6 @@ class Router: ): deployment.litellm_params.tpm = getattr(deployment, "tpm") - #### VALIDATE MODEL ######## - # check if model provider in supported providers - ( - _model, - custom_llm_provider, - dynamic_api_key, - api_base, - ) = litellm.get_llm_provider( - model=deployment.litellm_params.model, - custom_llm_provider=deployment.litellm_params.get( - "custom_llm_provider", None - ), - ) - # Check if user is trying to use model_name == "*" # this is a catch all model for their specific api key # if deployment.model_name == "*": @@ -4784,10 +4781,6 @@ class Router: env_name = params[param_key].replace("os.environ/", "") params[param_key] = os.environ.get(env_name, "") - # done reading model["litellm_params"] - if custom_llm_provider not in litellm.provider_list: - raise Exception(f"Unsupported provider - {custom_llm_provider}") - # # init OpenAI, Azure clients # InitalizeOpenAISDKClient.set_client( # litellm_router_instance=self, model=deployment.to_json(exclude_none=True) diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index ce3796bc3ba..8a15424d5be 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -659,3 +659,32 @@ def test_arouter_responses_api_bridge(): == "https://webhook.site/fba79dae-220a-4bb7-9a3a-8caa49604e55/openai/v1/responses?api-version=preview" ) assert mock_post.call_args.kwargs["json"]["model"] == "webinterface-o3-pro" + + +def test_add_invalid_provider_to_router(): + """ + Test that router.add_deployment raises an error if the provider is invalid + """ + from litellm.types.router import Deployment + + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-3.5-turbo", + "litellm_params": {"model": "gpt-3.5-turbo"}, + } + ], + ) + + with pytest.raises(Exception) as e: + router.add_deployment( + Deployment( + model_name="vertex_ai/*", + litellm_params={ + "model": "vertex_ai/*", + "custom_llm_provider": "vertex_ai_eu", + }, + ) + ) + + assert router.pattern_router.patterns == {}