mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
chore(proxy): tighten model-management authorization and stored-credential handling
The model-management write paths (POST /model/new, POST /model/update, PATCH /model/{id}/update) and POST /health/test_connection authorized callers at the whole-model level but applied no field-level checks and derived team scope from the request body rather than the persisted row. A team admin could rename a team model into the global routing pool by omitting model_info, zero out a model's pricing to escape budget tracking, bind another tenant's stored credential by name, point an existing model's api_base at an attacker host while keeping the encrypted key, reach internal/metadata endpoints via api_base, spoof a deployment id to delete a config-based model, and (via /health/test_connection) probe any deployment by id and exfiltrate its inherited key to a request-supplied api_base.
Team scope is now inherited from the persisted row when a PATCH omits it (so a team model keeps its internal name and cannot be reassigned to another team); pricing fields and litellm_credential_name binding are restricted to proxy admins; team-scoped models created by non-admins must supply their own credential; api_base/base_url are SSRF-validated on write under the existing user_url_validation toggle; a stored credential is cleared when the destination changes without a fresh key so it is never silently re-pointed; the deployment id is pinned to the row primary key; /model/new rejects an id that collides with a live deployment and /model/delete only evicts DB-origin router entries; team aliases are removed by their public name; and /health/test_connection authorizes against the resolved deployment's owner and refuses to send an inherited credential to a caller-overridden destination.
Adds unit coverage for each guard.
This commit is contained in:
parent
28c0d8579b
commit
3fdf9d9b1a
4 changed files with 1002 additions and 14 deletions
|
|
@ -78,6 +78,49 @@ def _reject_os_environ_references(params: dict) -> None:
|
|||
stack.append(value)
|
||||
|
||||
|
||||
_HEALTH_CREDENTIAL_FIELDS = (
|
||||
"api_key",
|
||||
"aws_secret_access_key",
|
||||
"aws_session_token",
|
||||
"vertex_credentials",
|
||||
)
|
||||
_HEALTH_DESTINATION_FIELDS = (
|
||||
"api_base",
|
||||
"base_url",
|
||||
"api_version",
|
||||
"vertex_location",
|
||||
"vertex_project",
|
||||
"aws_region_name",
|
||||
)
|
||||
|
||||
|
||||
def _reject_inherited_credential_redirect(
|
||||
config_litellm_params: dict, request_litellm_params: dict
|
||||
) -> None:
|
||||
"""Confused-deputy guard for /health/test_connection.
|
||||
|
||||
If a credential is inherited from the resolved deployment (present in the
|
||||
config params, not supplied by the request) while the request overrides the
|
||||
connection target, the inherited secret would be sent to a caller-chosen
|
||||
endpoint. Refuse it; the caller must supply their own credential or omit the
|
||||
destination override.
|
||||
"""
|
||||
inherited_credential = any(
|
||||
field in config_litellm_params and field not in request_litellm_params
|
||||
for field in _HEALTH_CREDENTIAL_FIELDS
|
||||
)
|
||||
overrides_destination = any(
|
||||
field in request_litellm_params for field in _HEALTH_DESTINATION_FIELDS
|
||||
)
|
||||
if inherited_credential and overrides_destination:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "Cannot override the connection target (e.g. api_base) while inheriting stored credentials. Supply your own api_key, or omit the destination override."
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def get_callback_identifier(callback):
|
||||
"""
|
||||
Get the callback identifier string, handling both strings and objects.
|
||||
|
|
@ -1672,7 +1715,7 @@ async def health_liveliness_options():
|
|||
tags=["health"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def test_model_connection(
|
||||
async def test_model_connection( # noqa: PLR0915
|
||||
request: Request,
|
||||
mode: Optional[
|
||||
Literal[
|
||||
|
|
@ -1747,6 +1790,7 @@ async def test_model_connection(
|
|||
dict: A dictionary containing the health check result with either success information or error details.
|
||||
"""
|
||||
from litellm.proxy._types import CommonProxyErrors
|
||||
from litellm.proxy.common_utils.resource_ownership import is_proxy_admin
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
ModelManagementAuthChecks,
|
||||
)
|
||||
|
|
@ -1766,11 +1810,23 @@ async def test_model_connection(
|
|||
# already resolved before reaching this endpoint; any remaining
|
||||
# reference must have come from the request body.
|
||||
_reject_os_environ_references(request_litellm_params)
|
||||
# Binding a named global credential resolves it by name with no ownership
|
||||
# model, so it is proxy-admin-only (consistent with model management).
|
||||
if request_litellm_params.get("litellm_credential_name") and not is_proxy_admin(
|
||||
user_api_key_dict
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={"error": "Only proxy admins can use litellm_credential_name."},
|
||||
)
|
||||
model_name = request_litellm_params.get("model")
|
||||
|
||||
# Look up model configuration from router if model name is provided
|
||||
# This gets the litellm_params from proxy config (with resolved env vars)
|
||||
config_litellm_params: dict = {}
|
||||
# model_info of the deployment resolved by id; used to authorize against
|
||||
# the deployment's real owner rather than the caller-supplied model_info.
|
||||
resolved_model_info: Optional[dict] = None
|
||||
if llm_router is not None:
|
||||
# Prefer disambiguation by deployment id (`model_info.id`) when
|
||||
# the caller supplies it. This is required when multiple
|
||||
|
|
@ -1793,6 +1849,9 @@ async def test_model_connection(
|
|||
config_litellm_params = deployment_by_id.litellm_params.model_dump(
|
||||
exclude_none=True
|
||||
)
|
||||
resolved_model_info = deployment_by_id.model_info.model_dump(
|
||||
exclude_none=True
|
||||
)
|
||||
elif model_name:
|
||||
# Fall back to model_name lookup for callers (e.g. the
|
||||
# "Add Model" wizard, or curl) that don't supply an id.
|
||||
|
|
@ -1829,12 +1888,25 @@ async def test_model_connection(
|
|||
# This allows users to override specific params while using config for credentials
|
||||
litellm_params = {**config_litellm_params, **request_litellm_params}
|
||||
|
||||
## Auth check
|
||||
# Refuse sending an inherited credential to a caller-overridden destination.
|
||||
_reject_inherited_credential_redirect(
|
||||
config_litellm_params=config_litellm_params,
|
||||
request_litellm_params=request_litellm_params,
|
||||
)
|
||||
|
||||
## Auth check — when the deployment was resolved by id, authorize against
|
||||
## its real owner (model_info), not the caller-supplied model_info, so a
|
||||
## team admin cannot probe another team's deployment by passing their own
|
||||
## team_id.
|
||||
await ModelManagementAuthChecks.can_user_make_model_call(
|
||||
model_params=Deployment(
|
||||
model_name="test_model",
|
||||
litellm_params=LiteLLM_Params(**litellm_params),
|
||||
model_info=model_info,
|
||||
model_info=(
|
||||
resolved_model_info
|
||||
if resolved_model_info is not None
|
||||
else model_info
|
||||
),
|
||||
),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
prisma_client=prisma_client,
|
||||
|
|
|
|||
|
|
@ -13,11 +13,12 @@ model/{model_id}/update - PATCH endpoint for model update.
|
|||
import asyncio
|
||||
import datetime
|
||||
import json
|
||||
from typing import Dict, List, Literal, Optional, Tuple, Union, cast
|
||||
from typing import Callable, Dict, List, Literal, Optional, Tuple, Union, cast
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.constants import LITELLM_PROXY_ADMIN_NAME
|
||||
|
|
@ -35,8 +36,13 @@ from litellm.proxy._types import (
|
|||
TeamModelDeleteRequest,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.litellm_core_utils.url_utils import SSRFError, validate_url
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
|
||||
decrypt_value_helper,
|
||||
encrypt_value_helper,
|
||||
)
|
||||
from litellm.proxy.common_utils.resource_ownership import is_proxy_admin
|
||||
from litellm.proxy.management_endpoints.common_utils import _is_user_team_admin
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
team_model_add,
|
||||
|
|
@ -62,6 +68,142 @@ from litellm.utils import get_utc_datetime
|
|||
router = APIRouter()
|
||||
|
||||
|
||||
# Credential fields that must never be silently re-pointed at a new endpoint,
|
||||
# and the routing fields that change where they are sent.
|
||||
_CREDENTIAL_LITELLM_PARAMS = (
|
||||
"api_key",
|
||||
"litellm_credential_name",
|
||||
"aws_secret_access_key",
|
||||
"aws_session_token",
|
||||
"vertex_credentials",
|
||||
)
|
||||
_DESTINATION_LITELLM_PARAMS = ("api_base", "base_url", "custom_llm_provider")
|
||||
|
||||
|
||||
def _field_explicitly_set(model: Optional[BaseModel], field: str) -> bool:
|
||||
"""True if `field` was explicitly provided (to a value or null) on a model."""
|
||||
return model is not None and field in model.model_fields_set
|
||||
|
||||
|
||||
def _validate_model_url_params(litellm_params: dict) -> None:
|
||||
"""SSRF-guard api_base/base_url on stored model configs, mirroring the
|
||||
request-time guard in auth_utils.is_request_body_safe so the proxy applies
|
||||
one consistent policy. Gated on litellm.user_url_validation (default True)
|
||||
with user_url_allowed_hosts as the escape hatch for internal endpoints."""
|
||||
if not getattr(litellm, "user_url_validation", False):
|
||||
return
|
||||
for url_field in ("api_base", "base_url"):
|
||||
url_value = litellm_params.get(url_field)
|
||||
if not url_value or not isinstance(url_value, str):
|
||||
continue
|
||||
try:
|
||||
validate_url(url_value)
|
||||
except SSRFError as e:
|
||||
raise ProxyException(
|
||||
message=(
|
||||
f"{url_field}={url_value!r} is rejected by the SSRF guard "
|
||||
f"({e}). Add the host to general_settings.user_url_allowed_hosts "
|
||||
"to allow it."
|
||||
),
|
||||
type=ProxyErrorTypes.validation_error.value,
|
||||
code=status.HTTP_400_BAD_REQUEST,
|
||||
param=url_field,
|
||||
)
|
||||
|
||||
|
||||
def _decrypted_param(value: object) -> object:
|
||||
"""Decrypt a stored litellm_param for comparison; non-string or
|
||||
non-encrypted values pass through unchanged."""
|
||||
return (
|
||||
decrypt_value_helper(value, "litellm_param", return_original_value=True)
|
||||
if isinstance(value, str)
|
||||
else value
|
||||
)
|
||||
|
||||
|
||||
def _strip_credentials_on_destination_change(
|
||||
merged_litellm_params: dict,
|
||||
patch_plaintext: dict,
|
||||
db_plaintext: "Callable[[str], object]",
|
||||
) -> None:
|
||||
"""If the patch changes a destination field (api_base/base_url/custom_llm_provider)
|
||||
without supplying a fresh credential, drop the inherited secret(s) from the
|
||||
merged params so a stored credential is never silently re-pointed at a new
|
||||
(possibly attacker-controlled) endpoint."""
|
||||
destination_changed = any(
|
||||
field in patch_plaintext and patch_plaintext[field] != db_plaintext(field)
|
||||
for field in _DESTINATION_LITELLM_PARAMS
|
||||
)
|
||||
supplied_new_credential = any(
|
||||
field in patch_plaintext for field in _CREDENTIAL_LITELLM_PARAMS
|
||||
)
|
||||
if destination_changed and not supplied_new_credential:
|
||||
for field in _CREDENTIAL_LITELLM_PARAMS:
|
||||
merged_litellm_params.pop(field, None)
|
||||
|
||||
|
||||
def _assert_privileged_model_fields_authorized(
|
||||
litellm_params: Optional[BaseModel],
|
||||
model_info: Optional[BaseModel],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> None:
|
||||
"""Pricing overrides and credential-name binding are proxy-admin-only.
|
||||
|
||||
Per-token/character pricing (SPECIAL_MODEL_INFO_PARAMS) gates spend/budget
|
||||
enforcement, so a non-admin must not set or clear it. litellm_credential_name
|
||||
resolves a globally-stored credential by name with no ownership model, so
|
||||
only a proxy admin may bind one to a model.
|
||||
"""
|
||||
if is_proxy_admin(user_api_key_dict):
|
||||
return
|
||||
for field in SPECIAL_MODEL_INFO_PARAMS:
|
||||
if _field_explicitly_set(litellm_params, field) or _field_explicitly_set(
|
||||
model_info, field
|
||||
):
|
||||
raise ProxyException(
|
||||
message=f"Only proxy admins can set model pricing fields (e.g. {field}).",
|
||||
type=ProxyErrorTypes.auth_error.value,
|
||||
code=status.HTTP_403_FORBIDDEN,
|
||||
param=field,
|
||||
)
|
||||
if (
|
||||
_field_explicitly_set(litellm_params, "litellm_credential_name")
|
||||
and litellm_params.litellm_credential_name is not None # type: ignore
|
||||
):
|
||||
raise ProxyException(
|
||||
message="Only proxy admins can bind a stored credential (litellm_credential_name) to a model.",
|
||||
type=ProxyErrorTypes.auth_error.value,
|
||||
code=status.HTTP_403_FORBIDDEN,
|
||||
param="litellm_credential_name",
|
||||
)
|
||||
|
||||
|
||||
def _assert_team_model_has_own_credential(
|
||||
model_params: Deployment, user_api_key_dict: UserAPIKeyAuth
|
||||
) -> None:
|
||||
"""A team-scoped model created by a non-admin must carry its own credential,
|
||||
so it cannot silently inherit the proxy's global provider keys at call time."""
|
||||
if is_proxy_admin(user_api_key_dict):
|
||||
return
|
||||
if model_params.model_info is None or model_params.model_info.team_id is None:
|
||||
return
|
||||
litellm_params = model_params.litellm_params
|
||||
has_credential = any(
|
||||
getattr(litellm_params, field, None)
|
||||
for field in ("api_key", "aws_secret_access_key", "vertex_credentials")
|
||||
)
|
||||
if not has_credential:
|
||||
raise ProxyException(
|
||||
message=(
|
||||
"Team-scoped models must provide their own credential (e.g. "
|
||||
"litellm_params.api_key) and cannot inherit the proxy's global keys."
|
||||
),
|
||||
type=ProxyErrorTypes.validation_error.value,
|
||||
code=status.HTTP_400_BAD_REQUEST,
|
||||
param="litellm_params.api_key",
|
||||
)
|
||||
|
||||
|
||||
async def update_team(*args, **kwargs):
|
||||
"""
|
||||
Backward-compatible shim for tests/legacy call sites that patch this symbol.
|
||||
|
|
@ -112,13 +254,16 @@ def update_db_model(
|
|||
merged_deployment_dict["model_name"] = updated_patch.model_name
|
||||
|
||||
# update litellm params
|
||||
patch_litellm_plaintext: dict = {}
|
||||
if updated_patch.litellm_params:
|
||||
patch_litellm_plaintext = updated_patch.litellm_params.model_dump(
|
||||
exclude_none=True
|
||||
)
|
||||
# SSRF-guard the patched destination (api_base/base_url).
|
||||
_validate_model_url_params(patch_litellm_plaintext)
|
||||
# Encrypt any sensitive values
|
||||
encrypted_params = {
|
||||
k: encrypt_value_helper(v)
|
||||
for k, v in updated_patch.litellm_params.model_dump(
|
||||
exclude_none=True
|
||||
).items()
|
||||
k: encrypt_value_helper(v) for k, v in patch_litellm_plaintext.items()
|
||||
}
|
||||
|
||||
merged_deployment_dict["litellm_params"].update(encrypted_params) # type: ignore
|
||||
|
|
@ -131,6 +276,13 @@ def update_db_model(
|
|||
updated_patch.model_info.model_dump(exclude_none=True)
|
||||
)
|
||||
|
||||
# The deployment id is the immutable primary key; never let a patched (or
|
||||
# freshly-constructed, auto-id'd) model_info blob change it.
|
||||
if db_model.model_info is not None and db_model.model_info.id is not None:
|
||||
merged_deployment_dict.setdefault("model_info", {})["id"] = ( # type: ignore
|
||||
db_model.model_info.id
|
||||
)
|
||||
|
||||
# Honor explicit-null clears LAST, after both merges, so a model_info blob the UI
|
||||
# passes through (which today re-sends the OLD pricing on every save) cannot
|
||||
# silently undo a litellm_params clear via .update().
|
||||
|
|
@ -157,6 +309,18 @@ def update_db_model(
|
|||
merged_deployment_dict["model_info"].pop(field, None) # type: ignore
|
||||
merged_deployment_dict.get("litellm_params", {}).pop(field, None) # type: ignore
|
||||
|
||||
# Refuse to silently re-point a stored credential at a new endpoint: if the
|
||||
# destination changed without a fresh credential, drop the inherited secret
|
||||
# so it must be re-entered rather than forwarded to the new destination.
|
||||
if updated_patch.litellm_params:
|
||||
_strip_credentials_on_destination_change(
|
||||
merged_deployment_dict["litellm_params"], # type: ignore
|
||||
patch_litellm_plaintext,
|
||||
lambda field: _decrypted_param(
|
||||
getattr(db_model.litellm_params, field, None)
|
||||
),
|
||||
)
|
||||
|
||||
# convert to prisma compatible format
|
||||
|
||||
prisma_compatible_model_dict = PrismaCompatibleUpdateDBModel()
|
||||
|
|
@ -274,6 +438,13 @@ async def patch_model(
|
|||
param="blocked",
|
||||
)
|
||||
|
||||
# Pricing overrides and credential binding are proxy-admin-only.
|
||||
_assert_privileged_model_fields_authorized(
|
||||
litellm_params=patch_data.litellm_params,
|
||||
model_info=patch_data.model_info,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
# Handle team model updates with proper alias management
|
||||
update_data = await _update_team_model_in_db(
|
||||
db_model=db_model,
|
||||
|
|
@ -443,9 +614,32 @@ async def _update_team_model_in_db(
|
|||
premium_user=premium_user,
|
||||
)
|
||||
|
||||
patch_team_id = patch_data.model_info.team_id if patch_data.model_info else None
|
||||
db_team_id = db_model.model_info.team_id if db_model.model_info else None
|
||||
explicit_patch_team_id = (
|
||||
patch_data.model_info.team_id if patch_data.model_info else None
|
||||
)
|
||||
|
||||
# No team_id in patch, proceed with standard update
|
||||
# A team admin must not move a model to a different team.
|
||||
if (
|
||||
explicit_patch_team_id is not None
|
||||
and db_team_id is not None
|
||||
and explicit_patch_team_id != db_team_id
|
||||
):
|
||||
raise ProxyException(
|
||||
message="Cannot reassign a model to a different team.",
|
||||
type=ProxyErrorTypes.auth_error.value,
|
||||
code=status.HTTP_403_FORBIDDEN,
|
||||
param="model_info.team_id",
|
||||
)
|
||||
|
||||
# Inherit the persisted team_id when the patch omits it, so a team-scoped
|
||||
# model keeps its team scoping (and its internal model_name) instead of being
|
||||
# renamed into the global pool via a model_info-less PATCH.
|
||||
patch_team_id = (
|
||||
explicit_patch_team_id if explicit_patch_team_id is not None else db_team_id
|
||||
)
|
||||
|
||||
# Genuinely non-team model: standard update.
|
||||
if patch_team_id is None:
|
||||
return update_db_model(db_model=db_model, updated_patch=patch_data)
|
||||
|
||||
|
|
@ -463,7 +657,6 @@ async def _update_team_model_in_db(
|
|||
patch_data.model_info.team_public_model_name = public_model_name
|
||||
|
||||
# Check if team assignment is new or changed
|
||||
db_team_id = db_model.model_info.team_id if db_model.model_info else None
|
||||
is_new_team_assignment = db_team_id != patch_team_id
|
||||
|
||||
if is_new_team_assignment:
|
||||
|
|
@ -826,7 +1019,13 @@ async def delete_model(
|
|||
# delete team model alias
|
||||
if model_params.model_info.team_id is not None:
|
||||
removed_model_aliases = await delete_team_model_alias(
|
||||
public_model_name=model_params.model_name,
|
||||
# The team alias is keyed on the PUBLIC model name; the internal
|
||||
# model_name is the unique UUID, which never matches an alias, so
|
||||
# using it here leaves the alias (and team.models entry) orphaned.
|
||||
public_model_name=(
|
||||
model_params.model_info.team_public_model_name
|
||||
or model_params.model_name
|
||||
),
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
|
|
@ -870,8 +1069,16 @@ async def delete_model(
|
|||
)
|
||||
|
||||
## DELETE FROM ROUTER ##
|
||||
# Only evict a router deployment that originated from the DB. A
|
||||
# config (static) deployment must never be removed by deleting a DB
|
||||
# row, even if a row shares its id.
|
||||
if llm_router is not None:
|
||||
llm_router.delete_deployment(id=model_info.id)
|
||||
router_deployment = llm_router.get_deployment(model_id=model_info.id)
|
||||
if (
|
||||
router_deployment is not None
|
||||
and router_deployment.model_info.db_model
|
||||
):
|
||||
llm_router.delete_deployment(id=model_info.id)
|
||||
|
||||
## CREATE AUDIT LOG ##
|
||||
asyncio.create_task(
|
||||
|
|
@ -1002,6 +1209,7 @@ async def add_new_model(
|
|||
"""
|
||||
from litellm.proxy.proxy_server import (
|
||||
general_settings,
|
||||
llm_router,
|
||||
premium_user,
|
||||
prisma_client,
|
||||
proxy_config,
|
||||
|
|
@ -1026,6 +1234,35 @@ async def add_new_model(
|
|||
premium_user=premium_user,
|
||||
)
|
||||
|
||||
# Pricing overrides and credential binding are proxy-admin-only.
|
||||
_assert_privileged_model_fields_authorized(
|
||||
litellm_params=model_params.litellm_params,
|
||||
model_info=model_params.model_info,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
# A non-admin team model must carry its own credential (no global-key fallback).
|
||||
_assert_team_model_has_own_credential(model_params, user_api_key_dict)
|
||||
# SSRF-guard the destination on create.
|
||||
_validate_model_url_params(
|
||||
model_params.litellm_params.model_dump(exclude_none=True)
|
||||
)
|
||||
# Reject a model_info.id that collides with a live deployment, so a DB
|
||||
# row cannot hijack (and on delete evict) a config/router deployment's id.
|
||||
if (
|
||||
model_params.model_info.id is not None
|
||||
and llm_router is not None
|
||||
and llm_router.has_model_id(model_params.model_info.id)
|
||||
):
|
||||
raise ProxyException(
|
||||
message=(
|
||||
f"model_info.id={model_params.model_info.id} already belongs to an "
|
||||
"existing deployment; omit model_info.id to auto-generate one."
|
||||
),
|
||||
type=ProxyErrorTypes.validation_error.value,
|
||||
code=status.HTTP_400_BAD_REQUEST,
|
||||
param="model_info.id",
|
||||
)
|
||||
|
||||
model_response: Optional[LiteLLM_ProxyModelTable] = None
|
||||
# update DB
|
||||
if store_model_in_db is True:
|
||||
|
|
@ -1192,6 +1429,13 @@ async def update_model(
|
|||
premium_user=premium_user,
|
||||
)
|
||||
|
||||
# Pricing/credential field authz, parity with patch_model / add_new_model.
|
||||
_assert_privileged_model_fields_authorized(
|
||||
litellm_params=model_params.litellm_params,
|
||||
model_info=model_params.model_info,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
# update DB
|
||||
if store_model_in_db is True:
|
||||
_existing_litellm_params_dict = dict(
|
||||
|
|
@ -1204,6 +1448,8 @@ async def update_model(
|
|||
_new_litellm_params_dict = model_params.litellm_params.dict(
|
||||
exclude_none=True
|
||||
)
|
||||
# SSRF-guard the patched destination on this legacy update path too.
|
||||
_validate_model_url_params(_new_litellm_params_dict)
|
||||
|
||||
### ENCRYPT PARAMS ###
|
||||
for k, v in _new_litellm_params_dict.items():
|
||||
|
|
@ -1225,6 +1471,15 @@ async def update_model(
|
|||
else:
|
||||
pass
|
||||
|
||||
# Don't silently re-point an inherited credential at a new endpoint.
|
||||
_strip_credentials_on_destination_change(
|
||||
merged_dictionary,
|
||||
_new_litellm_params_dict,
|
||||
lambda field: _decrypted_param(
|
||||
_existing_litellm_params_dict.get(field)
|
||||
),
|
||||
)
|
||||
|
||||
_data: dict = {
|
||||
"litellm_params": json.dumps(merged_dictionary), # type: ignore
|
||||
"updated_by": user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME,
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import contextlib
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
|
|
@ -1823,3 +1824,157 @@ def test_clean_endpoint_data_strips_credentials_keeps_routing_fields():
|
|||
assert "aws_access_key_id" not in cleaned
|
||||
assert cleaned.get("api_base") == "https://example.test/v1"
|
||||
assert cleaned.get("api_version") == "2024-10-21"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# VERIA-120: /health/test_connection confused-deputy hardening
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _health_test_connection_patches(mock_prisma, mock_router, mock_auth, mock_ahealth):
|
||||
def _noop_update(model_info, litellm_params):
|
||||
params = litellm_params.copy()
|
||||
params["messages"] = [{"role": "user", "content": "x"}]
|
||||
return params
|
||||
|
||||
return (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
|
||||
patch("litellm.proxy.proxy_server.llm_router", mock_router),
|
||||
patch("litellm.proxy.proxy_server.premium_user", True),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call",
|
||||
mock_auth,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.health_endpoints._health_endpoints.litellm.ahealth_check",
|
||||
mock_ahealth,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.health_endpoints._health_endpoints.run_with_timeout",
|
||||
AsyncMock(return_value={"status": "healthy"}),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.health_endpoints._health_endpoints._update_litellm_params_for_health_check",
|
||||
_noop_update,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.health_endpoints._health_endpoints._reject_os_environ_references",
|
||||
lambda params: None,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_test_connection_authorizes_against_resolved_deployment_owner():
|
||||
"""A team admin must not probe another team's deployment by id: the auth
|
||||
check must run against the RESOLVED deployment's team_id, not the
|
||||
caller-supplied model_info.team_id."""
|
||||
from litellm.types.router import Deployment, LiteLLM_Params
|
||||
|
||||
victim = Deployment(
|
||||
model_name="victim",
|
||||
litellm_params=LiteLLM_Params(
|
||||
model="azure/gpt-4o",
|
||||
api_key="sk-victim",
|
||||
api_base="https://victim.example/v1",
|
||||
),
|
||||
model_info={"id": "victim-id", "team_id": "team-OWNER"},
|
||||
)
|
||||
mock_router = MagicMock()
|
||||
mock_router.get_deployment.side_effect = lambda model_id: (
|
||||
victim if model_id == "victim-id" else None
|
||||
)
|
||||
|
||||
captured = {}
|
||||
|
||||
async def _capture(*, model_params, **kwargs):
|
||||
captured["team_id"] = model_params.model_info.team_id
|
||||
return True
|
||||
|
||||
mock_auth = AsyncMock(side_effect=_capture)
|
||||
mock_ahealth = AsyncMock(return_value={"status": "healthy"})
|
||||
|
||||
with contextlib.ExitStack() as stack:
|
||||
for p in _health_test_connection_patches(
|
||||
MagicMock(), mock_router, mock_auth, mock_ahealth
|
||||
):
|
||||
stack.enter_context(p)
|
||||
await health_test_model_connection(
|
||||
request=MagicMock(),
|
||||
mode="chat",
|
||||
litellm_params={"model": "azure/gpt-4o"}, # no destination override
|
||||
model_info={"id": "victim-id", "team_id": "team-ATTACKER"},
|
||||
user_api_key_dict=MagicMock(user_id="attacker", token="t"),
|
||||
)
|
||||
|
||||
assert captured.get("team_id") == "team-OWNER"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_test_connection_rejects_inherited_key_with_api_base_override():
|
||||
"""Refuse to send an inherited config api_key to a request-overridden
|
||||
api_base (the credential-exfiltration vector)."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.types.router import Deployment, LiteLLM_Params
|
||||
|
||||
victim = Deployment(
|
||||
model_name="victim",
|
||||
litellm_params=LiteLLM_Params(
|
||||
model="azure/gpt-4o",
|
||||
api_key="sk-victim",
|
||||
api_base="https://victim.example/v1",
|
||||
),
|
||||
model_info={"id": "victim-id", "team_id": "team-OWNER"},
|
||||
)
|
||||
mock_router = MagicMock()
|
||||
mock_router.get_deployment.side_effect = lambda model_id: (
|
||||
victim if model_id == "victim-id" else None
|
||||
)
|
||||
mock_auth = AsyncMock(return_value=True)
|
||||
mock_ahealth = AsyncMock(return_value={"status": "healthy"})
|
||||
|
||||
with contextlib.ExitStack() as stack:
|
||||
for p in _health_test_connection_patches(
|
||||
MagicMock(), mock_router, mock_auth, mock_ahealth
|
||||
):
|
||||
stack.enter_context(p)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await health_test_model_connection(
|
||||
request=MagicMock(),
|
||||
mode="chat",
|
||||
litellm_params={
|
||||
"model": "azure/gpt-4o",
|
||||
"api_base": "https://attacker.example/v1", # override, no api_key
|
||||
},
|
||||
model_info={"id": "victim-id"},
|
||||
user_api_key_dict=MagicMock(user_id="attacker", token="t"),
|
||||
)
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
mock_ahealth.assert_not_called() # the inherited key never went out
|
||||
|
||||
|
||||
def test_reject_inherited_credential_redirect_helper():
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy.health_endpoints._health_endpoints import (
|
||||
_reject_inherited_credential_redirect,
|
||||
)
|
||||
|
||||
# inherited credential + overridden destination -> rejected
|
||||
with pytest.raises(HTTPException):
|
||||
_reject_inherited_credential_redirect(
|
||||
config_litellm_params={"api_key": "sk-x", "api_base": "https://real"},
|
||||
request_litellm_params={"api_base": "https://attacker"},
|
||||
)
|
||||
# request supplies its own credential -> allowed
|
||||
_reject_inherited_credential_redirect(
|
||||
config_litellm_params={"api_key": "sk-x"},
|
||||
request_litellm_params={"api_key": "sk-own", "api_base": "https://attacker"},
|
||||
)
|
||||
# no destination override -> allowed
|
||||
_reject_inherited_credential_redirect(
|
||||
config_litellm_params={"api_key": "sk-x"},
|
||||
request_litellm_params={"model": "gpt-4o"},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1315,6 +1315,8 @@ class TestAddAndDeleteModelLifecycle:
|
|||
|
||||
mock_router = MagicMock()
|
||||
mock_router.delete_deployment = MagicMock()
|
||||
# The new model id is not already live in the router (no collision).
|
||||
mock_router.has_model_id = MagicMock(return_value=False)
|
||||
|
||||
_PS = "litellm.proxy.proxy_server"
|
||||
_ENCRYPT = "litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper"
|
||||
|
|
@ -1917,3 +1919,507 @@ class TestPatchModelBlockedAuthGate:
|
|||
)
|
||||
assert result is updated_row
|
||||
mock_prisma.db.litellm_proxymodeltable.update.assert_awaited_once()
|
||||
|
||||
|
||||
class TestModelMgmtAuthzHardening:
|
||||
"""Regression tests for the model-management authorization cluster.
|
||||
|
||||
Each test fails on the unpatched code path it targets and passes only with
|
||||
the corresponding guard in place.
|
||||
"""
|
||||
|
||||
# --- update_db_model: id pin + SSRF + credential re-point (115c / 111) ---
|
||||
|
||||
def _db_model_with_secret(self):
|
||||
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo
|
||||
|
||||
return Deployment(
|
||||
model_name="m",
|
||||
litellm_params=LiteLLM_Params(
|
||||
model="openai/gpt-4",
|
||||
api_key="sk-real",
|
||||
api_base="https://real.example.com",
|
||||
),
|
||||
model_info=ModelInfo(id="real-id"),
|
||||
)
|
||||
|
||||
def test_patch_cannot_change_stored_model_id(self):
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
update_db_model,
|
||||
)
|
||||
from litellm.types.router import ModelInfo
|
||||
|
||||
result = update_db_model(
|
||||
db_model=self._db_model_with_secret(),
|
||||
updated_patch=updateDeployment(model_info=ModelInfo(id="spoofed-id")),
|
||||
)
|
||||
info = json.loads(result["model_info"])
|
||||
assert info["id"] == "real-id"
|
||||
|
||||
def test_api_base_change_clears_inherited_credential(self):
|
||||
import litellm
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
update_db_model,
|
||||
)
|
||||
from litellm.types.router import updateLiteLLMParams
|
||||
|
||||
with (
|
||||
patch.object(litellm, "user_url_validation", False),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper",
|
||||
side_effect=lambda value, **kwargs: value,
|
||||
),
|
||||
):
|
||||
result = update_db_model(
|
||||
db_model=self._db_model_with_secret(),
|
||||
updated_patch=updateDeployment(
|
||||
litellm_params=updateLiteLLMParams(
|
||||
api_base="https://attacker.example.com"
|
||||
)
|
||||
),
|
||||
)
|
||||
params = json.loads(result["litellm_params"])
|
||||
assert "api_key" not in params # stored secret must not ride to new base
|
||||
|
||||
def test_api_base_change_with_new_key_keeps_credential(self):
|
||||
import litellm
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
update_db_model,
|
||||
)
|
||||
from litellm.types.router import updateLiteLLMParams
|
||||
|
||||
with (
|
||||
patch.object(litellm, "user_url_validation", False),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper",
|
||||
side_effect=lambda value, **kwargs: value,
|
||||
),
|
||||
):
|
||||
result = update_db_model(
|
||||
db_model=self._db_model_with_secret(),
|
||||
updated_patch=updateDeployment(
|
||||
litellm_params=updateLiteLLMParams(
|
||||
api_base="https://attacker.example.com", api_key="sk-new"
|
||||
)
|
||||
),
|
||||
)
|
||||
params = json.loads(result["litellm_params"])
|
||||
assert "api_key" in params # caller re-supplied a key, so it is kept
|
||||
|
||||
def test_unchanged_destination_keeps_credential(self):
|
||||
import litellm
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
update_db_model,
|
||||
)
|
||||
from litellm.types.router import updateLiteLLMParams
|
||||
|
||||
# Re-saving the same api_base (e.g. a UI resend) must NOT clear the key.
|
||||
with (
|
||||
patch.object(litellm, "user_url_validation", False),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper",
|
||||
side_effect=lambda value, **kwargs: value,
|
||||
),
|
||||
):
|
||||
result = update_db_model(
|
||||
db_model=self._db_model_with_secret(),
|
||||
updated_patch=updateDeployment(
|
||||
litellm_params=updateLiteLLMParams(
|
||||
api_base="https://real.example.com"
|
||||
)
|
||||
),
|
||||
)
|
||||
params = json.loads(result["litellm_params"])
|
||||
assert "api_key" in params
|
||||
|
||||
def test_validate_model_url_params_blocks_internal_ip(self):
|
||||
import litellm
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
_validate_model_url_params,
|
||||
)
|
||||
|
||||
with patch.object(litellm, "user_url_validation", True):
|
||||
with pytest.raises(ProxyException):
|
||||
_validate_model_url_params(
|
||||
{"api_base": "http://169.254.169.254/latest/meta-data/"}
|
||||
)
|
||||
# Honors the opt-out toggle (internal endpoints / Ollama).
|
||||
with (
|
||||
patch.object(litellm, "user_url_validation", False),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper",
|
||||
side_effect=lambda value, **kwargs: value,
|
||||
),
|
||||
):
|
||||
_validate_model_url_params({"api_base": "http://169.254.169.254/"})
|
||||
|
||||
# --- field-level authorization (173b / 115b / 173a) ---
|
||||
|
||||
def _non_admin(self):
|
||||
return UserAPIKeyAuth(
|
||||
user_id="team-admin", user_role=LitellmUserRoles.INTERNAL_USER
|
||||
)
|
||||
|
||||
def _admin(self):
|
||||
return UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
|
||||
def test_pricing_fields_are_proxy_admin_only(self):
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
_assert_privileged_model_fields_authorized,
|
||||
)
|
||||
from litellm.types.router import updateLiteLLMParams
|
||||
|
||||
# Non-admin setting pricing -> rejected.
|
||||
with pytest.raises(ProxyException) as e:
|
||||
_assert_privileged_model_fields_authorized(
|
||||
litellm_params=updateLiteLLMParams(input_cost_per_token=0.0),
|
||||
model_info=None,
|
||||
user_api_key_dict=self._non_admin(),
|
||||
)
|
||||
assert str(e.value.code) == "403"
|
||||
# Non-admin clearing pricing (explicit null) -> also rejected.
|
||||
with pytest.raises(ProxyException):
|
||||
_assert_privileged_model_fields_authorized(
|
||||
litellm_params=updateLiteLLMParams(output_cost_per_token=None),
|
||||
model_info=None,
|
||||
user_api_key_dict=self._non_admin(),
|
||||
)
|
||||
# Proxy admin -> allowed.
|
||||
_assert_privileged_model_fields_authorized(
|
||||
litellm_params=updateLiteLLMParams(input_cost_per_token=0.0),
|
||||
model_info=None,
|
||||
user_api_key_dict=self._admin(),
|
||||
)
|
||||
|
||||
def test_credential_name_binding_is_proxy_admin_only(self):
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
_assert_privileged_model_fields_authorized,
|
||||
)
|
||||
from litellm.types.router import updateLiteLLMParams
|
||||
|
||||
with pytest.raises(ProxyException) as e:
|
||||
_assert_privileged_model_fields_authorized(
|
||||
litellm_params=updateLiteLLMParams(
|
||||
litellm_credential_name="someone-elses-cred"
|
||||
),
|
||||
model_info=None,
|
||||
user_api_key_dict=self._non_admin(),
|
||||
)
|
||||
assert str(e.value.code) == "403"
|
||||
# Admin may bind a stored credential.
|
||||
_assert_privileged_model_fields_authorized(
|
||||
litellm_params=updateLiteLLMParams(litellm_credential_name="shared-cred"),
|
||||
model_info=None,
|
||||
user_api_key_dict=self._admin(),
|
||||
)
|
||||
|
||||
def test_team_model_requires_own_credential(self):
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
_assert_team_model_has_own_credential,
|
||||
)
|
||||
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo
|
||||
|
||||
keyless_team_model = Deployment(
|
||||
model_name="team-gpt",
|
||||
litellm_params=LiteLLM_Params(model="openai/gpt-4"),
|
||||
model_info=ModelInfo(team_id="team-x"),
|
||||
)
|
||||
with pytest.raises(ProxyException) as e:
|
||||
_assert_team_model_has_own_credential(keyless_team_model, self._non_admin())
|
||||
assert str(e.value.code) == "400"
|
||||
# With its own key -> allowed.
|
||||
keyed_team_model = Deployment(
|
||||
model_name="team-gpt",
|
||||
litellm_params=LiteLLM_Params(model="openai/gpt-4", api_key="sk-team"),
|
||||
model_info=ModelInfo(team_id="team-x"),
|
||||
)
|
||||
_assert_team_model_has_own_credential(keyed_team_model, self._non_admin())
|
||||
# Proxy admin may create a key-less team model (uses configured keys).
|
||||
_assert_team_model_has_own_credential(keyless_team_model, self._admin())
|
||||
|
||||
# --- team-scope inheritance: cannot rename a team model into the global pool (115a / 168a) ---
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_model_cannot_be_renamed_to_global_via_omitted_model_info(self):
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
_update_team_model_in_db,
|
||||
)
|
||||
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo
|
||||
|
||||
db_model = Deployment(
|
||||
model_name="model_name_teamX_abc123",
|
||||
litellm_params=LiteLLM_Params(model="azure/gpt-4o", api_key="sk-team"),
|
||||
model_info=ModelInfo(
|
||||
team_id="teamX", team_public_model_name="my-team-alias"
|
||||
),
|
||||
)
|
||||
# Attacker PATCHes only model_name (omits model_info) to a global name.
|
||||
patch_data = updateDeployment(model_name="gpt-4")
|
||||
admin_of_team = UserAPIKeyAuth(
|
||||
user_id="team-admin", user_role=LitellmUserRoles.INTERNAL_USER
|
||||
)
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.premium_user", True),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints.team_model_add"
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints.team_model_delete"
|
||||
),
|
||||
):
|
||||
result = await _update_team_model_in_db(
|
||||
db_model=db_model,
|
||||
patch_data=patch_data,
|
||||
user_api_key_dict=admin_of_team,
|
||||
prisma_client=MockPrismaClient(team_exists=True), # type: ignore
|
||||
)
|
||||
|
||||
# The internal team model_name must be preserved, not overwritten to "gpt-4".
|
||||
assert result.get("model_name", "model_name_teamX_abc123").startswith(
|
||||
"model_name_teamX_"
|
||||
)
|
||||
assert result.get("model_name") != "gpt-4"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_patch_rejects_moving_model_to_a_different_team(self):
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
_update_team_model_in_db,
|
||||
)
|
||||
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo
|
||||
|
||||
db_model = Deployment(
|
||||
model_name="model_name_teamX_abc123",
|
||||
litellm_params=LiteLLM_Params(model="azure/gpt-4o", api_key="sk-team"),
|
||||
model_info=ModelInfo(team_id="teamX"),
|
||||
)
|
||||
patch_data = updateDeployment(model_info=ModelInfo(team_id="teamY"))
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.premium_user", True),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints.team_model_add"
|
||||
),
|
||||
):
|
||||
with pytest.raises(ProxyException) as e:
|
||||
await _update_team_model_in_db(
|
||||
db_model=db_model,
|
||||
patch_data=patch_data,
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_id="team-admin",
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
),
|
||||
prisma_client=MockPrismaClient(team_exists=True), # type: ignore
|
||||
)
|
||||
assert str(e.value.code) == "403"
|
||||
|
||||
# --- model_id spoofing + alias residue (126 / 168b) ---
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_new_model_rejects_colliding_model_id(self):
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
add_new_model,
|
||||
)
|
||||
from litellm.proxy.proxy_server import ProxyException
|
||||
from litellm.types.router import Deployment, LiteLLM_Params
|
||||
|
||||
mock_router = MagicMock()
|
||||
# The supplied id already belongs to a live (config) deployment.
|
||||
mock_router.has_model_id = MagicMock(return_value=True)
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db = MagicMock()
|
||||
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
|
||||
|
||||
_PS = "litellm.proxy.proxy_server"
|
||||
with (
|
||||
patch(f"{_PS}.prisma_client", mock_prisma),
|
||||
patch(f"{_PS}.store_model_in_db", True),
|
||||
patch(f"{_PS}.proxy_config", MagicMock()),
|
||||
patch(f"{_PS}.proxy_logging_obj", MagicMock()),
|
||||
patch(f"{_PS}.general_settings", {}),
|
||||
patch(f"{_PS}.premium_user", True),
|
||||
patch(f"{_PS}.llm_router", mock_router),
|
||||
):
|
||||
with pytest.raises(ProxyException) as e:
|
||||
await add_new_model(
|
||||
model_params=Deployment(
|
||||
model_name="x",
|
||||
litellm_params=LiteLLM_Params(
|
||||
model="openai/gpt-4", api_key="k"
|
||||
),
|
||||
model_info={"id": "config-model-1"},
|
||||
),
|
||||
user_api_key_dict=self._admin(),
|
||||
)
|
||||
assert str(e.value.code) == "400"
|
||||
mock_prisma.db.litellm_proxymodeltable.create.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_preserves_config_origin_router_deployment(self):
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
ModelInfoDelete,
|
||||
delete_model as delete_model_endpoint,
|
||||
)
|
||||
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo
|
||||
|
||||
model_id = "config-model-1"
|
||||
db_row = LiteLLM_ProxyModelTable(
|
||||
model_id=model_id,
|
||||
model_name="m",
|
||||
litellm_params={"model": "openai/gpt-4"},
|
||||
model_info={"id": model_id}, # no team_id (global/non-team row)
|
||||
created_by="a",
|
||||
updated_by="a",
|
||||
)
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db = MagicMock()
|
||||
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
|
||||
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(
|
||||
return_value=db_row
|
||||
)
|
||||
mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row)
|
||||
|
||||
mock_router = MagicMock()
|
||||
mock_router.delete_deployment = MagicMock()
|
||||
# The live router entry for this id is config-origin (db_model unset).
|
||||
mock_router.get_deployment = MagicMock(
|
||||
return_value=Deployment(
|
||||
model_name="m",
|
||||
litellm_params=LiteLLM_Params(model="openai/gpt-4"),
|
||||
model_info=ModelInfo(id=model_id),
|
||||
)
|
||||
)
|
||||
|
||||
_PS = "litellm.proxy.proxy_server"
|
||||
with (
|
||||
patch(f"{_PS}.prisma_client", mock_prisma),
|
||||
patch(f"{_PS}.store_model_in_db", True),
|
||||
patch(f"{_PS}.premium_user", True),
|
||||
patch(f"{_PS}.llm_router", mock_router),
|
||||
):
|
||||
await delete_model_endpoint(
|
||||
model_info=ModelInfoDelete(id=model_id),
|
||||
user_api_key_dict=self._admin(),
|
||||
)
|
||||
mock_prisma.db.litellm_proxymodeltable.delete.assert_awaited_once()
|
||||
# The config deployment must NOT be evicted from the router.
|
||||
mock_router.delete_deployment.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_uses_team_public_model_name_for_alias(self):
|
||||
from litellm.proxy.management_endpoints import (
|
||||
model_management_endpoints as mod,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
ModelInfoDelete,
|
||||
delete_model as delete_model_endpoint,
|
||||
)
|
||||
|
||||
model_id = "team-dep-1"
|
||||
db_row = LiteLLM_ProxyModelTable(
|
||||
model_id=model_id,
|
||||
model_name="model_name_teamX_uuid", # internal UUID name
|
||||
litellm_params={"model": "openai/gpt-4"},
|
||||
model_info={
|
||||
"id": model_id,
|
||||
"team_id": "teamX",
|
||||
"team_public_model_name": "team-alias",
|
||||
},
|
||||
created_by="a",
|
||||
updated_by="a",
|
||||
)
|
||||
team_row = LiteLLM_TeamTable(
|
||||
team_id="teamX",
|
||||
team_alias="t",
|
||||
models=["team-alias"],
|
||||
members_with_roles=[],
|
||||
)
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db = MagicMock()
|
||||
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
|
||||
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(
|
||||
return_value=db_row
|
||||
)
|
||||
mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row)
|
||||
mock_prisma.db.litellm_teamtable = AsyncMock()
|
||||
mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row)
|
||||
mock_prisma.db.litellm_teamtable.update = AsyncMock()
|
||||
|
||||
mock_router = MagicMock()
|
||||
mock_router.get_deployment = MagicMock(return_value=None)
|
||||
|
||||
captured = {}
|
||||
|
||||
async def fake_delete_alias(public_model_name, prisma_client):
|
||||
captured["public_model_name"] = public_model_name
|
||||
return []
|
||||
|
||||
_PS = "litellm.proxy.proxy_server"
|
||||
with (
|
||||
patch(f"{_PS}.prisma_client", mock_prisma),
|
||||
patch(f"{_PS}.store_model_in_db", True),
|
||||
patch(f"{_PS}.premium_user", True),
|
||||
patch(f"{_PS}.llm_router", mock_router),
|
||||
patch.object(mod, "delete_team_model_alias", side_effect=fake_delete_alias),
|
||||
):
|
||||
await delete_model_endpoint(
|
||||
model_info=ModelInfoDelete(id=model_id),
|
||||
user_api_key_dict=self._admin(),
|
||||
)
|
||||
# Must use the public alias, not the internal UUID model_name.
|
||||
assert captured.get("public_model_name") == "team-alias"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_legacy_update_model_pricing_is_proxy_admin_only(self):
|
||||
"""The legacy POST /model/update path must enforce the same pricing gate
|
||||
as PATCH, so it can't be used to bypass it."""
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
update_model,
|
||||
)
|
||||
from litellm.proxy.proxy_server import ProxyException
|
||||
from litellm.types.router import (
|
||||
ModelInfo,
|
||||
updateDeployment,
|
||||
updateLiteLLMParams,
|
||||
)
|
||||
|
||||
model_id = "db-model-1"
|
||||
existing_row = MagicMock()
|
||||
existing_row.litellm_params = {"model": "openai/gpt-4o-mini", "api_key": "sk-x"}
|
||||
existing_row.model_dump.return_value = {
|
||||
"model_name": "m",
|
||||
"litellm_params": existing_row.litellm_params,
|
||||
"model_info": {"id": model_id, "team_id": "teamX"},
|
||||
}
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(
|
||||
return_value=existing_row
|
||||
)
|
||||
mock_prisma.db.litellm_proxymodeltable.update = AsyncMock()
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
|
||||
patch("litellm.proxy.proxy_server.llm_router", MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.store_model_in_db", True),
|
||||
patch("litellm.proxy.proxy_server.premium_user", True),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call",
|
||||
new=AsyncMock(return_value=None),
|
||||
),
|
||||
):
|
||||
with pytest.raises(ProxyException) as e:
|
||||
await update_model(
|
||||
model_params=updateDeployment(
|
||||
litellm_params=updateLiteLLMParams(input_cost_per_token=0.0),
|
||||
model_info=ModelInfo(id=model_id),
|
||||
),
|
||||
user_api_key_dict=self._non_admin(),
|
||||
)
|
||||
assert str(e.value.code) == "403"
|
||||
mock_prisma.db.litellm_proxymodeltable.update.assert_not_called()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue