mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
a34206f67e
commit
ea6b4b08d3
6 changed files with 120 additions and 27 deletions
|
|
@ -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",
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue