fix: skip cascade cleanup when other deployments share the model name

When deleting a global model, check whether another deployment still
uses the same model_name before removing it from team model lists.
This prevents breaking load-balanced setups where multiple deployments
share a single model_name.
This commit is contained in:
Jonathan Wrede 2026-05-10 20:26:25 +00:00
parent 13480a9c3a
commit c7921679aa
2 changed files with 79 additions and 10 deletions

View file

@ -810,19 +810,30 @@ async def delete_model(
# Remove deleted global model from all teams that reference it by name.
# Only for non-team-scoped models; team-scoped cleanup is handled above.
# Skip if another deployment still uses the same model_name (load-balancing).
if model_params.model_info.team_id is None:
try:
affected_teams = await prisma_client.db.litellm_teamtable.find_many(
where={"models": {"has": model_params.model_name}}
)
for team in affected_teams:
updated_models = [
m for m in team.models if m != model_params.model_name
]
await prisma_client.db.litellm_teamtable.update(
where={"team_id": team.team_id},
data={"models": updated_models},
other_deployments = (
await prisma_client.db.litellm_proxymodeltable.find_many(
where={
"model_name": model_params.model_name,
"model_id": {"not": model_info.id},
},
take=1,
)
)
if not other_deployments:
affected_teams = await prisma_client.db.litellm_teamtable.find_many(
where={"models": {"has": model_params.model_name}}
)
for team in affected_teams:
updated_models = [
m for m in team.models if m != model_params.model_name
]
await prisma_client.db.litellm_teamtable.update(
where={"team_id": team.team_id},
data={"models": updated_models},
)
except Exception as e:
verbose_proxy_logger.warning(
"Failed to remove model %s from teams: %s",

View file

@ -1401,6 +1401,7 @@ class TestDeleteModelCascadeTeams:
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(
return_value=db_row
)
mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row)
mock_prisma.db.litellm_teamtable = AsyncMock()
mock_prisma.db.litellm_teamtable.find_many = AsyncMock(
@ -1500,6 +1501,63 @@ class TestDeleteModelCascadeTeams:
# The cascade find_many should NOT be called for team-scoped models
mock_prisma.db.litellm_teamtable.find_many.assert_not_called()
@pytest.mark.asyncio
async def test_delete_skips_cascade_when_other_deployment_exists(self):
"""When another deployment shares the same model_name, don't remove from teams."""
from litellm.proxy.management_endpoints.model_management_endpoints import (
delete_model as delete_model_endpoint,
ModelInfoDelete,
)
model_id = "deployment-1"
admin_user = UserAPIKeyAuth(
user_id="test-admin", user_role=LitellmUserRoles.PROXY_ADMIN
)
db_row = LiteLLM_ProxyModelTable(
model_id=model_id,
model_name="gpt-4",
litellm_params={"model": "openai/gpt-4", "api_key": "key-1"},
model_info={"id": model_id},
created_by="test-admin",
updated_by="test-admin",
)
other_deployment = MagicMock()
other_deployment.model_id = "deployment-2"
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.find_many = AsyncMock(
return_value=[other_deployment]
)
mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row)
mock_prisma.db.litellm_teamtable = AsyncMock()
mock_prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_teamtable.update = 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_logging_obj", MagicMock()),
patch(f"{_PS}.general_settings", {}),
patch(f"{_PS}.premium_user", True),
patch(f"{_PS}.llm_router", MagicMock()),
):
await delete_model_endpoint(
model_info=ModelInfoDelete(id=model_id),
user_api_key_dict=admin_user,
)
# Should not query teams since another deployment exists
mock_prisma.db.litellm_teamtable.find_many.assert_not_called()
mock_prisma.db.litellm_teamtable.update.assert_not_called()
class TestGetTeamDeployments:
"""Tests for _get_team_deployments which filters by model_name prefix + Python-side team_id check."""