mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
(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 <ishaanjaffer0324@gmail.com>
This commit is contained in:
parent
56b8e857bd
commit
962c33ee72
2 changed files with 46 additions and 24 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 == {}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue