move to use_prisma_migrate by default + resolve team-only models on auth checks + UI - add sagemaker on UI (#13117)

* fix(proxy_cli.py): make use_prisma_migrate proxy default

Fixes https://github.com/BerriAI/litellm/issues/13046

 Prisma migrate deploy prevents resetting db

* fix(auth_checks.py): resolve team only models while doing auth checks on model access groups

Fixes issue where key had access via an access group, but team only model could not be called

* test(test_router.py): add unit testing

* feat(provider_specific_fields.tsx): add aws sagemaker on UI
This commit is contained in:
Krish Dholakia 2025-07-29 21:56:18 -07:00 • committed by GitHub
parent a34206f67e
commit ea6b4b08d3
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 120 additions and 27 deletions

View file

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

View file

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

View file

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

View file

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

View file

@ -243,6 +243,29 @@ const PROVIDER_CREDENTIAL_FIELDS: Record<Providers, ProviderCredentialField[]> =
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",

View file

@ -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<string, string> = {
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<string, string> = {
[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;