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,
+ )