Ensure disable_llm_api_endpoints works + Add wildcard model support for 'team-byok' model (#13278)

* fix(route_checks.py): ensure disable llm api endpoints is correctly set

* fix(route_checks.py): raise httpexception

raise expected exceptions

* fix(router.py): handle team only wildcard models

fixes issue where team only wildcard models were not considered during auth checks

* fix(router.py): handle team only wildcard models

fixes issue where team only wildcard models were not considered during auth checks
This commit is contained in:
Krish Dholakia 2025-08-04 23:19:51 -07:00 • committed by GitHub
parent af67c3576f
commit eb49f987de
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 121 additions and 44 deletions

View file

@ -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

File diff suppressed because one or more lines are too long

View file

@ -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"]
api_key: os.environ/OPENAI_API_KEY

View file

@ -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,

View file

@ -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:
"""

View file

@ -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)
)

View file

@ -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

View file

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