mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
fix(model management): only proxy admins may set auto router compression fields
This commit is contained in:
parent
8544faec91
commit
1bcfdaf7cb
2 changed files with 180 additions and 0 deletions
|
|
@ -850,6 +850,12 @@ async def patch_model(
|
|||
existing_litellm_params=db_model.litellm_params,
|
||||
)
|
||||
|
||||
ModelManagementAuthChecks.can_user_set_auto_router_compression(
|
||||
litellm_params=patch_data.litellm_params,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
existing_litellm_params=db_model.litellm_params,
|
||||
)
|
||||
|
||||
_raise_on_strategy_router_write_violation(
|
||||
incoming_params=patch_data.litellm_params,
|
||||
existing_params=db_model.litellm_params,
|
||||
|
|
@ -1603,6 +1609,17 @@ async def _update_existing_team_model_assignment(
|
|||
# No team_model_add/delete calls required; public name is already registered
|
||||
|
||||
|
||||
def _stored_auto_router_compression(
|
||||
litellm_params: GenericLiteLLMParams | None,
|
||||
) -> tuple[str | None, str | None]:
|
||||
"""The stored auto router compression values, decrypting if necessary."""
|
||||
if litellm_params is None:
|
||||
return (None, None)
|
||||
routing: Final = litellm_params.auto_router_routing_compression
|
||||
model: Final = litellm_params.auto_router_model_compression
|
||||
return (routing, model)
|
||||
|
||||
|
||||
class ModelManagementAuthChecks:
|
||||
"""
|
||||
Common auth checks for model management endpoints
|
||||
|
|
@ -1657,6 +1674,34 @@ class ModelManagementAuthChecks:
|
|||
param="litellm_credential_name",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def can_user_set_auto_router_compression(
|
||||
litellm_params: GenericLiteLLMParams | None,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
existing_litellm_params: GenericLiteLLMParams | None = None,
|
||||
) -> Literal[True]:
|
||||
if litellm_params is None:
|
||||
return True
|
||||
incoming: Final = (
|
||||
litellm_params.auto_router_routing_compression,
|
||||
litellm_params.auto_router_model_compression,
|
||||
)
|
||||
if incoming == (None, None):
|
||||
return True
|
||||
if incoming == _stored_auto_router_compression(existing_litellm_params):
|
||||
return True
|
||||
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN:
|
||||
return True
|
||||
raise ProxyException(
|
||||
message=(
|
||||
"Only a proxy admin can set auto_router_routing_compression / auto_router_model_compression on a "
|
||||
f"model. Your role={user_api_key_dict.user_role}."
|
||||
),
|
||||
type=ProxyErrorTypes.auth_error.value,
|
||||
code=status.HTTP_403_FORBIDDEN,
|
||||
param="auto_router_model_compression",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def allow_team_model_action(
|
||||
model_params: Deployment | updateDeployment,
|
||||
|
|
@ -1990,6 +2035,11 @@ async def add_new_model(
|
|||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
ModelManagementAuthChecks.can_user_set_auto_router_compression(
|
||||
litellm_params=model_params.litellm_params,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
_raise_on_strategy_router_write_violation(
|
||||
incoming_params=model_params.litellm_params,
|
||||
existing_params=None,
|
||||
|
|
@ -2174,6 +2224,12 @@ async def update_model(
|
|||
existing_litellm_params=deployment.litellm_params,
|
||||
)
|
||||
|
||||
ModelManagementAuthChecks.can_user_set_auto_router_compression(
|
||||
litellm_params=model_params.litellm_params,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
existing_litellm_params=deployment.litellm_params,
|
||||
)
|
||||
|
||||
_raise_on_strategy_router_write_violation(
|
||||
incoming_params=model_params.litellm_params,
|
||||
existing_params=deployment.litellm_params,
|
||||
|
|
|
|||
|
|
@ -400,6 +400,130 @@ class TestModelManagementAuthChecks:
|
|||
)
|
||||
assert exc_info.value.code == "403"
|
||||
|
||||
def test_can_user_set_auto_router_compression_admin_success(self):
|
||||
result = ModelManagementAuthChecks.can_user_set_auto_router_compression(
|
||||
litellm_params=LiteLLM_Params(model="auto_router/test", auto_router_model_compression="none"),
|
||||
user_api_key_dict=self.admin_user,
|
||||
)
|
||||
assert result is True
|
||||
|
||||
def test_can_user_set_auto_router_compression_without_fields_allows_any_role(self):
|
||||
result = ModelManagementAuthChecks.can_user_set_auto_router_compression(
|
||||
litellm_params=LiteLLM_Params(model="test_model"),
|
||||
user_api_key_dict=self.team_admin_user,
|
||||
)
|
||||
assert result is True
|
||||
|
||||
def test_can_user_set_auto_router_compression_team_admin_fails(self):
|
||||
with pytest.raises(Exception, match="Only a proxy admin can set auto_router") as exc_info:
|
||||
ModelManagementAuthChecks.can_user_set_auto_router_compression(
|
||||
litellm_params=LiteLLM_Params(model="auto_router/test", auto_router_model_compression="none"),
|
||||
user_api_key_dict=self.team_admin_user,
|
||||
)
|
||||
assert exc_info.value.code == "403"
|
||||
|
||||
def test_can_user_set_auto_router_compression_routing_only_team_admin_fails(self):
|
||||
with pytest.raises(Exception, match="Only a proxy admin can set auto_router") as exc_info:
|
||||
ModelManagementAuthChecks.can_user_set_auto_router_compression(
|
||||
litellm_params=LiteLLM_Params(model="auto_router/test", auto_router_routing_compression="headroom"),
|
||||
user_api_key_dict=self.team_admin_user,
|
||||
)
|
||||
assert exc_info.value.code == "403"
|
||||
|
||||
def test_can_user_set_auto_router_compression_unchanged_existing_allows_any_role(self):
|
||||
result = ModelManagementAuthChecks.can_user_set_auto_router_compression(
|
||||
litellm_params=LiteLLM_Params(model="auto_router/test", auto_router_model_compression="none"),
|
||||
user_api_key_dict=self.team_admin_user,
|
||||
existing_litellm_params=LiteLLM_Params(model="auto_router/test", auto_router_model_compression="none"),
|
||||
)
|
||||
assert result is True
|
||||
|
||||
def test_can_user_set_auto_router_compression_changed_model_field_team_admin_fails(self):
|
||||
with pytest.raises(Exception, match="Only a proxy admin can set auto_router") as exc_info:
|
||||
ModelManagementAuthChecks.can_user_set_auto_router_compression(
|
||||
litellm_params=LiteLLM_Params(model="auto_router/test", auto_router_model_compression="headroom"),
|
||||
user_api_key_dict=self.team_admin_user,
|
||||
existing_litellm_params=LiteLLM_Params(model="auto_router/test", auto_router_model_compression="none"),
|
||||
)
|
||||
assert exc_info.value.code == "403"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_new_model_rejects_auto_router_compression_for_non_admin(self):
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
add_new_model,
|
||||
)
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), # test-quality-ok: endpoint reads proxy server globals with no injection seam
|
||||
patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: endpoint reads proxy server globals with no injection seam
|
||||
patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: endpoint reads proxy server globals with no injection seam
|
||||
patch( # test-quality-ok: prior auth check needs a live DB; only the compression check is under test
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call",
|
||||
new=AsyncMock(return_value=None),
|
||||
),
|
||||
):
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await add_new_model(
|
||||
model_params=Deployment(
|
||||
model_name="compression-model",
|
||||
litellm_params=LiteLLM_Params(
|
||||
model="auto_router/complexity_router",
|
||||
auto_router_model_compression="none",
|
||||
),
|
||||
model_info={"id": "compression-create-test"},
|
||||
),
|
||||
user_api_key_dict=self.team_admin_user,
|
||||
)
|
||||
assert exc_info.value.code == "403"
|
||||
mock_prisma.db.litellm_proxymodeltable.create.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_patch_model_rejects_auto_router_compression_for_non_admin(self):
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
patch_model,
|
||||
)
|
||||
from litellm.types.router import updateLiteLLMParams
|
||||
|
||||
model_id = "compression-patch-test"
|
||||
db_model = Deployment(
|
||||
model_name="compression-model",
|
||||
litellm_params=LiteLLM_Params(model="auto_router/complexity_router"),
|
||||
model_info={"id": model_id},
|
||||
)
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), # test-quality-ok: endpoint reads proxy server globals with no injection seam
|
||||
patch("litellm.proxy.proxy_server.llm_router", MagicMock()), # test-quality-ok: endpoint reads proxy server globals with no injection seam
|
||||
patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: endpoint reads proxy server globals with no injection seam
|
||||
patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: endpoint reads proxy server globals with no injection seam
|
||||
patch( # test-quality-ok: stubs the DB row fetch; only the compression check is under test
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints.get_db_model",
|
||||
new=AsyncMock(return_value=db_model),
|
||||
),
|
||||
patch( # test-quality-ok: prior auth check needs a live DB; only the compression check is under test
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call",
|
||||
new=AsyncMock(return_value=None),
|
||||
),
|
||||
patch( # test-quality-ok: asserts the DB write is never reached on rejection
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints._update_team_model_in_db",
|
||||
new=AsyncMock(),
|
||||
) as mock_update,
|
||||
):
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await patch_model(
|
||||
model_id=model_id,
|
||||
patch_data=updateDeployment(
|
||||
litellm_params=updateLiteLLMParams(
|
||||
auto_router_model_compression="none",
|
||||
)
|
||||
),
|
||||
user_api_key_dict=self.team_admin_user,
|
||||
)
|
||||
assert exc_info.value.code == "403"
|
||||
mock_update.assert_not_awaited()
|
||||
|
||||
|
||||
class MockModelTable:
|
||||
def __init__(self, model_aliases: Dict[str, str], include: Optional[dict] = None):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue