diff --git a/enterprise/litellm_enterprise/proxy/auth/route_checks.py b/enterprise/litellm_enterprise/proxy/auth/route_checks.py index 1d4bfc664d5..6cce781faf3 100644 --- a/enterprise/litellm_enterprise/proxy/auth/route_checks.py +++ b/enterprise/litellm_enterprise/proxy/auth/route_checks.py @@ -20,7 +20,6 @@ class EnterpriseRouteChecks: status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail=f"🚨🚨🚨 DISABLING LLM API ENDPOINTS is an Enterprise feature\n🚨 {CommonProxyErrors.not_premium_user.value}", ) - return False return get_secret_bool("DISABLE_LLM_API_ENDPOINTS") is True diff --git a/litellm/proxy/_experimental/out/model_hub_table.html b/litellm/proxy/_experimental/out/model_hub_table/index.html similarity index 100% rename from litellm/proxy/_experimental/out/model_hub_table.html rename to litellm/proxy/_experimental/out/model_hub_table/index.html diff --git a/litellm/proxy/_experimental/out/onboarding.html b/litellm/proxy/_experimental/out/onboarding.html deleted file mode 100644 index f30076137a2..00000000000 --- a/litellm/proxy/_experimental/out/onboarding.html +++ /dev/null @@ -1 +0,0 @@ -LiteLLM Dashboard \ No newline at end of file diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index ad4dabf9d04..b587913a08a 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -2,27 +2,4 @@ model_list: - model_name: openai-test litellm_params: model: gpt-3.5-turbo - api_key: os.environ/OPENAI_API_KEY - -guardrails: - - guardrail_name: azure-text-moderation - litellm_params: - guardrail: azure/text_moderations - mode: "post_call" - api_key: os.environ/AZURE_GUARDRAIL_API_KEY - api_base: os.environ/AZURE_GUARDRAIL_API_BASE - -prompts: - - prompt_id: test_my_json_prompt - litellm_params: - prompt_integration: dotprompt - prompt_id: test_hello_world_prompt - prompt_file: /Users/krrishdholakia/Documents/litellm/litellm/proxy/test_prompts/test_hello_world_prompt.prompt - - prompt_id: test_hello_world_prompt_2 - litellm_params: - prompt_integration: dotprompt - prompt_id: test_hello_world_prompt_2 - prompt_file: /Users/krrishdholakia/Documents/litellm/litellm/proxy/test_prompts/test_hello_world_prompt.prompt - -litellm_settings: - callbacks: ["datadog_llm_observability"] \ No newline at end of file + api_key: os.environ/OPENAI_API_KEY \ No newline at end of file diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 32ba2080a36..89c012c20ae 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -288,9 +288,6 @@ def _is_api_route_allowed( if valid_token is None: raise Exception("Invalid proxy server token passed. valid_token=None.") - # Check if management routes are disabled and raise exception if they are - RouteChecks.should_call_route(route=route, valid_token=valid_token) - if not _is_user_proxy_admin(user_obj=user_obj): # if non-admin RouteChecks.non_proxy_admin_allowed_routes_check( user_obj=user_obj, diff --git a/litellm/proxy/auth/route_checks.py b/litellm/proxy/auth/route_checks.py index cd158fd70a7..e4d4f321159 100644 --- a/litellm/proxy/auth/route_checks.py +++ b/litellm/proxy/auth/route_checks.py @@ -25,6 +25,8 @@ class RouteChecks: from litellm_enterprise.proxy.auth.route_checks import EnterpriseRouteChecks EnterpriseRouteChecks.should_call_route(route=route) + except HTTPException as e: + raise e except Exception: pass @@ -386,7 +388,7 @@ class RouteChecks: if "thread" in request.url.path or "assistant" in request.url.path: return True return False - + @staticmethod def is_generate_content_route(route: str) -> bool: """ diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 31395fe7da0..00711afee70 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -49,6 +49,7 @@ from litellm.proxy.auth.auth_utils import ( from litellm.proxy.auth.handle_jwt import JWTAuthManager, JWTHandler from litellm.proxy.auth.oauth2_check import check_oauth2_token from litellm.proxy.auth.oauth2_proxy_hook import handle_oauth2_proxy_request +from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.common_utils.http_parsing_utils import ( _read_request_body, _safe_get_request_headers, @@ -259,6 +260,7 @@ def get_api_key( from litellm.proxy.common_utils.http_parsing_utils import ( _safe_get_request_query_params, ) + api_key = api_key passed_in_key: Optional[str] = None if isinstance(custom_litellm_key_header, str): @@ -279,7 +281,11 @@ def get_api_key( elif isinstance(azure_apim_header, str): passed_in_key = azure_apim_header api_key = azure_apim_header - elif RouteChecks.is_generate_content_route(route=route) and request is not None and _safe_get_request_query_params(request).get("key"): + elif ( + RouteChecks.is_generate_content_route(route=route) + and request is not None + and _safe_get_request_query_params(request).get("key") + ): google_auth_key: str = _safe_get_request_query_params(request).get("key") or "" passed_in_key = google_auth_key api_key = google_auth_key @@ -1141,6 +1147,8 @@ async def user_api_key_auth( request_data = await _read_request_body(request=request) route: str = get_request_route(request=request) + ## CHECK IF ROUTE IS ALLOWED + user_api_key_auth_obj = await _user_api_key_auth_builder( request=request, api_key=api_key, @@ -1152,6 +1160,9 @@ async def user_api_key_auth( custom_litellm_key_header=custom_litellm_key_header, ) + ## ENSURE DISABLE ROUTE WORKS ACROSS ALL USER AUTH FLOWS ## + RouteChecks.should_call_route(route=route, valid_token=user_api_key_auth_obj) + end_user_id = get_end_user_id_from_request_body( request_data, _safe_get_request_headers(request) ) diff --git a/litellm/router.py b/litellm/router.py index ca724006fdf..b89815877b5 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -406,6 +406,9 @@ class Router: self.default_max_parallel_requests = default_max_parallel_requests self.provider_default_deployment_ids: List[str] = [] self.pattern_router = PatternMatchRouter() + self.team_pattern_routers: Dict[str, PatternMatchRouter] = ( + {} + ) # {"TEAM_ID": PatternMatchRouter} self.auto_routers: Dict[str, "AutoRouter"] = {} if model_list is not None: @@ -1207,7 +1210,7 @@ class Router: verbose_router_logger.error( f"Fallback also failed: {fallback_error}" ) - raise fallback_error + raise fallback_error return FallbackStreamWrapper(stream_with_fallbacks()) @@ -5109,6 +5112,19 @@ class Router: if deployment.model_info.id: self.provider_default_deployment_ids.append(deployment.model_info.id) + _team_id = deployment.model_info.get("team_id") + _team_public_model_name = deployment.model_info.get("team_public_model_name") + if ( + _team_id is not None + and _team_public_model_name is not None + and "*" in _team_public_model_name + ): + if _team_id not in self.team_pattern_routers: + self.team_pattern_routers[_team_id] = PatternMatchRouter() + self.team_pattern_routers[_team_id].add_pattern( + _team_public_model_name, deployment.to_json(exclude_none=True) + ) + # Azure GPT-Vision Enhancements, users can pass os.environ/ data_sources = deployment.litellm_params.get("dataSources", []) or [] @@ -5920,19 +5936,17 @@ class Router: Map a team model name to a team-specific model name. Returns: - - team_model_name: str - the team-specific model name + - deployment id: str - the deployment id of the team-specific model - None: if no team-specific model name is found """ - for model in self.model_list: - model_team_id = model["model_info"].get("team_id") - model_team_public_model_name = model["model_info"].get( - "team_public_model_name" - ) - if ( - model_team_id == team_id - and model_team_public_model_name == team_model_name - ): - return model["model_name"] + models = self.get_model_list(model_name=team_model_name, team_id=team_id) + if not models: + return None + for model in models: + if model.get("model_info", {}).get("team_id") == team_id: + return model.get("model_name") + + ## wildcard models return None def should_include_deployment( @@ -6073,6 +6087,7 @@ class Router: if team_id specified, returns matching team-specific models """ + if hasattr(self, "model_list"): returned_models: List[DeploymentTypedDict] = [] @@ -6087,7 +6102,17 @@ class Router: ) if len(returned_models) == 0: # check if wildcard route - potential_wildcard_models = self.pattern_router.route(model_name) + potential_wildcard_models = self.pattern_router.route(model_name) or [] + + ## check for team-specific wildcard models + if team_id is not None and team_id in self.team_pattern_routers: + potential_team_only_wildcard_models = ( + self.team_pattern_routers[team_id].route(model_name) or [] + ) + potential_wildcard_models.extend( + potential_team_only_wildcard_models + ) + if model_name is not None and potential_wildcard_models is not None: for m in potential_wildcard_models: deployment_typed_dict = DeploymentTypedDict(**m) # type: ignore @@ -6519,6 +6544,7 @@ class Router: messages: Optional[List[Dict[str, str]]] = None, input: Optional[Union[str, List]] = None, specific_deployment: Optional[bool] = False, + request_kwargs: Optional[Dict] = None, ) -> Tuple[str, Union[List, Dict]]: """ Common checks for 'get_available_deployment' across sync + async call. @@ -6530,6 +6556,14 @@ class Router: - List, if multiple models chosen - Dict, if specific model chosen """ + + request_team_id: Optional[str] = None + if request_kwargs is not None: + metadata = request_kwargs.get("metadata") or {} + litellm_metadata = request_kwargs.get("litellm_metadata") or {} + request_team_id = metadata.get( + "user_api_key_team_id" + ) or litellm_metadata.get("user_api_key_team_id") # check if aliases set on litellm model alias map if specific_deployment is True: return model, self._get_deployment_by_litellm_model(model=model) @@ -6552,9 +6586,22 @@ class Router: pattern_deployments = self.pattern_router.get_deployments_by_pattern( model=model, ) + if pattern_deployments: return model, pattern_deployments + if ( + request_team_id is not None + and request_team_id in self.team_pattern_routers + ): + pattern_deployments = self.team_pattern_routers[ + request_team_id + ].get_deployments_by_pattern( + model=model, + ) + if pattern_deployments: + return model, pattern_deployments + # check if default deployment is set if self.default_deployment is not None: updated_deployment = copy.deepcopy( @@ -6622,6 +6669,7 @@ class Router: messages=messages, input=input, specific_deployment=specific_deployment, + request_kwargs=request_kwargs, ) # type: ignore # IF TEAM ID SPECIFIED ON MODEL, AND REQUEST CONTAINS USER_API_KEY_TEAM_ID, FILTER OUT MODELS THAT ARE NOT IN THE TEAM diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index e8b6567a97d..0a6ea31a1ba 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -1357,3 +1357,47 @@ async def test_async_function_with_fallbacks_common_utils(): args=(), kwargs={}, # No model key ) + + +def test_should_include_deployment(): + """Test that Router.should_include_deployment returns the correct response""" + router = litellm.Router( + model_list=[ + { + "model_name": "model_name_a28a12f9-3e44-4861-bd4f-325f2d309ce8_cd5dc6fb-b046-4e05-ae1d-32ba4d936266", + "litellm_params": {"model": "openai/*"}, + "model_info": { + "team_id": "a28a12f9-3e44-4861-bd4f-325f2d309ce8", + "team_public_model_name": "openai/*", + }, + } + ], + ) + + model = { + "model_name": "model_name_a28a12f9-3e44-4861-bd4f-325f2d309ce8_cd5dc6fb-b046-4e05-ae1d-32ba4d936266", + "litellm_params": { + "api_key": "sk-proj-1234567890", + "custom_llm_provider": "openai", + "use_in_pass_through": False, + "use_litellm_proxy": False, + "merge_reasoning_content_in_choices": False, + "model": "openai/*", + }, + "model_info": { + "id": "95f58039-d54a-4d1c-b700-5e32e99a1120", + "db_model": True, + "updated_by": "64a2f787-0863-4d76-9516-2dc49c1598e8", + "created_by": "64a2f787-0863-4d76-9516-2dc49c1598e8", + "team_id": "a28a12f9-3e44-4861-bd4f-325f2d309ce8", + "team_public_model_name": "openai/*", + "mode": "completion", + "access_groups": ["restricted-models-openai"], + }, + } + model_name = "openai/o4-mini-deep-research" + team_id = "a28a12f9-3e44-4861-bd4f-325f2d309ce8" + assert router.get_model_list( + model_name=model_name, + team_id=team_id, + )