mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
af67c3576f
commit
eb49f987de
9 changed files with 121 additions and 44 deletions
|
|
@ -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
|
|
@ -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
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue