diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 1e57e5326a1..32ba2080a36 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -1152,6 +1152,7 @@ def _can_object_call_model( llm_router: Optional[Router], models: List[str], team_model_aliases: Optional[Dict[str, str]] = None, + team_id: Optional[str] = None, object_type: Literal["user", "team", "key", "org"] = "user", fallback_depth: int = 0, ) -> Literal[True]: @@ -1184,6 +1185,7 @@ def _can_object_call_model( llm_router=llm_router, models=models, team_model_aliases=team_model_aliases, + team_id=team_id, object_type=object_type, fallback_depth=fallback_depth + 1, ) @@ -1202,7 +1204,9 @@ def _can_object_call_model( access_groups: Dict[str, List[str]] = defaultdict(list) if llm_router: - access_groups = llm_router.get_model_access_groups(model_name=model) + access_groups = llm_router.get_model_access_groups( + model_name=model, team_id=team_id + ) if ( len(access_groups) > 0 and llm_router is not None @@ -1288,6 +1292,7 @@ async def can_key_call_model( llm_router=llm_router, models=valid_token.models, team_model_aliases=valid_token.team_model_aliases, + team_id=valid_token.team_id, object_type="key", ) @@ -1326,6 +1331,7 @@ def can_team_access_model( llm_router=llm_router, models=team_object.models if team_object else [], team_model_aliases=team_model_aliases, + team_id=team_object.team_id if team_object else None, object_type="team", ) diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index 729db54b1b0..c18c14fa920 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -301,6 +301,7 @@ class ProxyInitializationHelpers: return None # Let uvicorn choose the default loop on Windows return "uvloop" + @click.command() @click.option( "--host", default="0.0.0.0", help="Host for the server to listen on.", envvar="HOST" @@ -467,7 +468,7 @@ class ProxyInitializationHelpers: @click.option( "--use_prisma_migrate", is_flag=True, - default=False, + default=True, help="Use prisma migrate instead of prisma db push for database schema updates", ) @click.option("--local", is_flag=True, default=False, help="for local debugging") @@ -789,13 +790,12 @@ def run_server( # noqa: PLR0915 # DO NOT DELETE - enables global variables to work across files from litellm.proxy.proxy_server import app # noqa - + # --- SEPARATE HEALTH APP LOGIC --- # To run the health app separately, use: # uvicorn litellm.proxy.health_app_factory:build_health_app --factory --host 0.0.0.0 --port=4001 # This is compatible with the SEPARATE_HEALTH_APP Docker/supervisord pattern. # --- END SEPARATE HEALTH APP LOGIC --- - # Skip server startup if requested (after all setup is done) if skip_server_startup: print( # noqa diff --git a/litellm/router.py b/litellm/router.py index 9bb8fcf6997..d02392c59f2 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -4701,8 +4701,6 @@ class Router: model_id: _model_info, } ) - - ## OLD MODEL REGISTRATION ## Kept to prevent breaking changes _model_name = deployment.litellm_params.model @@ -4741,7 +4739,7 @@ class Router: return None else: raise e - + def _is_auto_router_deployment(self, litellm_params: LiteLLM_Params) -> bool: """ Check if the deployment is an auto-router deployment. @@ -4751,7 +4749,7 @@ class Router: if litellm_params.model.startswith("auto_router/"): return True return False - + def init_auto_router_deployment(self, deployment: Deployment): """ Initialize the auto-router deployment. @@ -4759,18 +4757,31 @@ class Router: This will initialize the auto-router and add it to the auto-routers dictionary. """ from litellm.router_strategy.auto_router.auto_router import AutoRouter - auto_router_config_path: Optional[str] = deployment.litellm_params.auto_router_config_path + + auto_router_config_path: Optional[str] = ( + deployment.litellm_params.auto_router_config_path + ) auto_router_config: Optional[str] = deployment.litellm_params.auto_router_config if auto_router_config_path is None and auto_router_config is None: - raise ValueError("auto_router_config_path or auto_router_config is required for auto-router deployments. Please set it in the litellm_params") - - default_model: Optional[str] = deployment.litellm_params.auto_router_default_model + raise ValueError( + "auto_router_config_path or auto_router_config is required for auto-router deployments. Please set it in the litellm_params" + ) + + default_model: Optional[str] = ( + deployment.litellm_params.auto_router_default_model + ) if default_model is None: - raise ValueError("auto_router_default_model is required for auto-router deployments. Please set it in the litellm_params") - - embedding_model: Optional[str] = deployment.litellm_params.auto_router_embedding_model + raise ValueError( + "auto_router_default_model is required for auto-router deployments. Please set it in the litellm_params" + ) + + embedding_model: Optional[str] = ( + deployment.litellm_params.auto_router_embedding_model + ) if embedding_model is None: - raise ValueError("auto_router_embedding_model is required for auto-router deployments. Please set it in the litellm_params") + raise ValueError( + "auto_router_embedding_model is required for auto-router deployments. Please set it in the litellm_params" + ) autor_router: AutoRouter = AutoRouter( model_name=deployment.model_name, @@ -4781,7 +4792,9 @@ class Router: litellm_router_instance=self, ) if deployment.model_name in self.auto_routers: - raise ValueError(f"Auto-router deployment {deployment.model_name} already exists. Please use a different model name.") + raise ValueError( + f"Auto-router deployment {deployment.model_name} already exists. Please use a different model name." + ) self.auto_routers[deployment.model_name] = autor_router def deployment_is_active_for_environment(self, deployment: Deployment) -> bool: @@ -5777,6 +5790,8 @@ class Router: Return all deployments of a model name Used for accurate 'get_model_list'. + + if team_id specified, only return team-specific models """ returned_models: List[DeploymentTypedDict] = [] for model in self.model_list: @@ -5917,7 +5932,10 @@ class Router: return None def get_model_access_groups( - self, model_name: Optional[str] = None, model_access_group: Optional[str] = None + self, + model_name: Optional[str] = None, + model_access_group: Optional[str] = None, + team_id: Optional[str] = None, ) -> Dict[str, List[str]]: """ If model_name is provided, only return access groups for that model. @@ -5925,12 +5943,13 @@ class Router: Parameters: - model_name: Optional[str] - the received model name from the user (can be a wildcard route). If set, will only return access groups for that model. - model_access_group: Optional[str] - the received model access group from the user. If set, will only return models for that access group. + - team_id: Optional[str] - the team id, to resolve team-specific models """ from collections import defaultdict access_groups = defaultdict(list) - model_list = self.get_model_list(model_name=model_name) + model_list = self.get_model_list(model_name=model_name, team_id=team_id) if model_list: for m in model_list: _model_info = m.get("model_info") @@ -6518,7 +6537,7 @@ class Router: ) try: parent_otel_span = _get_parent_otel_span_from_kwargs(request_kwargs) - + ######################################################### # Execute Pre-Routing Hooks # this hook can modify the model, messages before the routing decision is made @@ -6535,8 +6554,6 @@ class Router: messages = pre_routing_hook_response.messages ######################################################### - - healthy_deployments = await self.async_get_healthy_deployments( model=model, request_kwargs=request_kwargs, @@ -6671,12 +6688,9 @@ class Router: input=input, specific_deployment=specific_deployment, ) - + return None - - - def get_available_deployment( self, model: str, diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 638dfd774bf..cacdd51ee90 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -1035,3 +1035,32 @@ def test_router_apply_default_settings(): assert ( has_forward_headers_check ), "Expected ForwardClientSideHeadersByModelGroup to be added to callbacks" + + +def test_router_get_model_access_groups_team_only_models(): + """ + Test that Router.get_model_access_groups returns the correct response for team-only models + """ + router = litellm.Router( + model_list=[ + { + "model_name": "my-custom-model-name", + "litellm_params": {"model": "gpt-3.5-turbo"}, + "model_info": { + "team_id": "team_1", + "access_groups": ["default-models"], + "team_public_model_name": "gpt-3.5-turbo", + }, + }, + ] + ) + + access_groups = router.get_model_access_groups( + model_name="gpt-3.5-turbo", team_id=None + ) + assert len(access_groups) == 0 + + access_groups = router.get_model_access_groups( + model_name="gpt-3.5-turbo", team_id="team_1" + ) + assert list(access_groups.keys()) == ["default-models"] diff --git a/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx b/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx index 3c6784bf85a..d537f5c4612 100644 --- a/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx +++ b/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx @@ -243,6 +243,29 @@ const PROVIDER_CREDENTIAL_FIELDS: Record = tooltip: "You can provide the raw key or the environment variable (e.g. `os.environ/MY_SECRET_KEY`)." } ], + [Providers.SageMaker]: [ + { + key: "aws_access_key_id", + label: "AWS Access Key ID", + type: "password", + required: false, + tooltip: "You can provide the raw key or the environment variable (e.g. `os.environ/MY_SECRET_KEY`)." + }, + { + key: "aws_secret_access_key", + label: "AWS Secret Access Key", + type: "password", + required: false, + tooltip: "You can provide the raw key or the environment variable (e.g. `os.environ/MY_SECRET_KEY`)." + }, + { + key: "aws_region_name", + label: "AWS Region Name", + placeholder: "us-east-1", + required: false, + tooltip: "You can provide the raw key or the environment variable (e.g. `os.environ/MY_SECRET_KEY`)." + } + ], [Providers.Ollama]: [ { key: "api_base", diff --git a/ui/litellm-dashboard/src/components/provider_info_helpers.tsx b/ui/litellm-dashboard/src/components/provider_info_helpers.tsx index 0e02a92770a..ed638995bdd 100644 --- a/ui/litellm-dashboard/src/components/provider_info_helpers.tsx +++ b/ui/litellm-dashboard/src/components/provider_info_helpers.tsx @@ -12,6 +12,7 @@ export enum Providers { Vertex_AI = "Vertex AI (Anthropic, Gemini, etc.)", Google_AI_Studio = "Google AI Studio", Bedrock = "Amazon Bedrock", + SageMaker = "AWS SageMaker", Groq = "Groq", MistralAI = "Mistral AI", Deepseek = "Deepseek", @@ -59,7 +60,8 @@ export const provider_map: Record = { FireworksAI: "fireworks_ai", Triton: "triton", Deepgram: "deepgram", - ElevenLabs: "elevenlabs" + ElevenLabs: "elevenlabs", + SageMaker: "sagemaker_chat" }; const asset_logos_folder = '/ui/assets/logos/'; @@ -70,6 +72,7 @@ export const providerLogoMap: Record = { [Providers.Azure]: `${asset_logos_folder}microsoft_azure.svg`, [Providers.Azure_AI_Studio]: `${asset_logos_folder}microsoft_azure.svg`, [Providers.Bedrock]: `${asset_logos_folder}bedrock.svg`, + [Providers.SageMaker]: `${asset_logos_folder}bedrock.svg`, [Providers.Cerebras]: `${asset_logos_folder}cerebras.svg`, [Providers.Cohere]: `${asset_logos_folder}cohere.svg`, [Providers.Databricks]: `${asset_logos_folder}databricks.svg`, @@ -129,6 +132,8 @@ export const getPlaceholder = (selectedProvider: string): string => { return "claude-3-opus"; } else if (selectedProvider == Providers.Bedrock) { return "claude-3-opus"; + } else if (selectedProvider == Providers.SageMaker) { + return "sagemaker/jumpstart-dft-meta-textgeneration-llama-2-7b"; } else if (selectedProvider == Providers.Google_AI_Studio) { return "gemini-pro"; } else if (selectedProvider == Providers.Azure_AI_Studio) { @@ -176,6 +181,22 @@ export const getPlaceholder = (selectedProvider: string): string => { } }); } + + // Special case for sagemaker + // we need both sagemaker and sagemaker_chat models to show on dropdown + if (providerKey == Providers.SageMaker) { + console.log("Adding sagemaker chat models"); + Object.entries(modelMap).forEach(([key, value]) => { + if ( + value !== null && + typeof value === "object" && + "litellm_provider" in (value as object) && + ((value as any)["litellm_provider"] === "sagemaker_chat") + ) { + providerModels.push(key); + } + }); + } } return providerModels;