Merge pull request #41694 from BerriAI/litellm_rename_model_sync_allowlists

fix(proxy): propagate db model renames to key, team, org, project and user model allowlists
This commit is contained in:
ryan-crabbe-berri 2026-09-17 17:42:06 -07:00 • committed by GitHub
commit 85fe646776
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 264 additions and 13 deletions

View file

@ -88,6 +88,7 @@ from litellm.proxy.management_helpers.auto_router_permissions import (
authorize_member_auto_router_team,
authorize_member_auto_router_write,
)
from litellm.proxy.management_helpers.model_allowlist_rename_sync import sync_model_allowlists_for_renamed_model
from litellm.proxy.spend_tracking.ptu_feature_flag import (
PTU_COST_ATTRIBUTION_ENV_VAR,
is_ptu_cost_attribution_enabled,
@ -986,6 +987,7 @@ async def patch_model(
premium_user,
prisma_client,
store_model_in_db,
user_api_key_cache,
)
try:
@ -1134,6 +1136,14 @@ async def patch_model(
new_name=stored_model_name,
llm_router=llm_router,
)
await sync_model_allowlists_for_renamed_model(
prisma_client=prisma_client,
model_id=model_id,
old_name=db_model.model_name,
new_name=stored_model_name,
llm_router=llm_router,
user_api_key_cache=user_api_key_cache,
)
# Clear cache and reload models (uses config setting or defaults to preserving config models for DB updates)
live_before_reload: Final = live_model_ids_snapshot()
@ -2435,6 +2445,7 @@ async def update_model(
premium_user,
prisma_client,
store_model_in_db,
user_api_key_cache,
)
try:
@ -2568,6 +2579,14 @@ async def update_model(
new_name=renamed_to,
llm_router=llm_router,
)
await sync_model_allowlists_for_renamed_model(
prisma_client=prisma_client,
model_id=_model_id,
old_name=deployment.model_name,
new_name=renamed_to,
llm_router=llm_router,
user_api_key_cache=user_api_key_cache,
)
# Clear cache and reload models (uses config setting or defaults to preserving config models for DB updates)
live_before_reload: Final = live_model_ids_snapshot()

View file

@ -24,7 +24,7 @@ class _DeploymentCountRow(BaseModel):
deployment_count: int
class _RawExecutor(Protocol):
class RawExecutor(Protocol):
async def query_raw(self, query: str, *args: str) -> Sequence[object]: ...
@ -54,7 +54,7 @@ _REMOVE_MODEL_NAME_SQL: Final = (
)
def _raw_executor(prisma_client: object) -> _RawExecutor:
def raw_executor(prisma_client: object) -> RawExecutor:
db: Final = AccessGroupRepository(prisma_client).prisma_client.db # pyright: ignore[reportAny] # untyped Prisma client
return writer_wrapper(db) # pyright: ignore[reportAny, reportReturnType] # untyped Prisma client behind the pin
@ -75,14 +75,14 @@ def _served_by_a_config_deployment(llm_router: Router | None, model_name: str, m
)
async def _still_backed(executor: _RawExecutor, llm_router: Router | None, model_name: str, model_id: str) -> bool:
async def still_backed(executor: RawExecutor, llm_router: Router | None, model_name: str, model_id: str) -> bool:
if _served_by_a_config_deployment(llm_router, model_name, model_id):
return True
count_rows: Final = await executor.query_raw(_BACKING_DEPLOYMENTS_SQL, model_name)
return any(_DeploymentCountRow.model_validate(row).deployment_count > 0 for row in count_rows)
async def _rewrite_groups(executor: _RawExecutor, sql: str, *names: str) -> None:
async def _rewrite_groups(executor: RawExecutor, sql: str, *names: str) -> None:
touched_rows: Final = await executor.query_raw(sql, *names)
await invalidate_access_group_caches(
tuple(_TouchedGroupRow.model_validate(row).access_group_id for row in touched_rows)
@ -99,8 +99,8 @@ async def sync_access_groups_for_renamed_model(
) -> None:
if old_name == new_name:
return
executor: Final = _raw_executor(prisma_client)
old_name_still_backed: Final = await _still_backed(executor, llm_router, old_name, model_id)
executor: Final = raw_executor(prisma_client)
old_name_still_backed: Final = await still_backed(executor, llm_router, old_name, model_id)
await _rewrite_groups(
executor, _APPEND_MODEL_NAME_SQL if old_name_still_backed else _REPLACE_MODEL_NAME_SQL, old_name, new_name
)
@ -113,7 +113,7 @@ async def sync_access_groups_for_deleted_model(
model_name: str,
llm_router: Router | None,
) -> None:
executor: Final = _raw_executor(prisma_client)
if await _still_backed(executor, llm_router, model_name, model_id):
executor: Final = raw_executor(prisma_client)
if await still_backed(executor, llm_router, model_name, model_id):
return
await _rewrite_groups(executor, _REMOVE_MODEL_NAME_SQL, model_name)

View file

@ -0,0 +1,109 @@
"""
Keep the `models` allowlists on keys, teams, organizations, projects and users pointing at
deployment names that still exist.
Those allowlists store public model names, not ids, so a deployment rename that leaves them
alone denies the new name while the old entry grants a name nothing serves any more.
"""
from collections.abc import Callable
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final
from pydantic import BaseModel
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.management_helpers.access_group_model_sync import raw_executor, still_backed
from litellm.router import Router
class _TouchedRow(BaseModel):
kind: str
object_id: str
team_alias: str | None = None
@dataclass(frozen=True, slots=True)
class _AllowlistTable:
kind: str
table: str
id_column: str
cache_keys: Callable[[_TouchedRow], tuple[str, ...]]
alias_column: str | None = None
def update_cte(self, set_clause: str, where_clause: str) -> str:
alias: Final = f'"{self.alias_column}"' if self.alias_column else "NULL::text"
return (
f'{self.kind}_rows AS (UPDATE "{self.table}" SET "models" = {set_clause} WHERE {where_clause} '
f"RETURNING '{self.kind}' AS kind, \"{self.id_column}\" AS object_id, {alias} AS team_alias)"
)
def _team_cache_keys(row: _TouchedRow) -> tuple[str, ...]:
return (f"team_id:{row.object_id}", *((f"team_alias:{row.team_alias}",) if row.team_alias else ()))
def _key_cache_keys(row: _TouchedRow) -> tuple[str, ...]:
return (row.object_id,)
def _org_cache_keys(row: _TouchedRow) -> tuple[str, ...]:
return (f"org_id:{row.object_id}", f"org_id:{row.object_id}:with_budget")
def _project_cache_keys(row: _TouchedRow) -> tuple[str, ...]:
return (f"project_id:{row.object_id}",)
def _user_cache_keys(row: _TouchedRow) -> tuple[str, ...]:
return (row.object_id,)
_ALLOWLIST_TABLES: Final = (
_AllowlistTable("team", "LiteLLM_TeamTable", "team_id", _team_cache_keys, alias_column="team_alias"),
_AllowlistTable("key", "LiteLLM_VerificationToken", "token", _key_cache_keys),
_AllowlistTable("org", "LiteLLM_OrganizationTable", "organization_id", _org_cache_keys),
_AllowlistTable("project", "LiteLLM_ProjectTable", "project_id", _project_cache_keys),
_AllowlistTable("user", "LiteLLM_UserTable", "user_id", _user_cache_keys),
)
_CACHE_KEYS_BY_KIND: Final = MappingProxyType({table.kind: table.cache_keys for table in _ALLOWLIST_TABLES})
def _rewrite_sql(set_clause: str, where_clause: str) -> str:
"""One statement touching every allowlist table, so the rewrite lands everywhere or nowhere."""
ctes: Final = ", ".join(table.update_cte(set_clause, where_clause) for table in _ALLOWLIST_TABLES)
rows: Final = " UNION ALL ".join(
f"SELECT kind, object_id, team_alias FROM {table.kind}_rows" for table in _ALLOWLIST_TABLES
)
return f"WITH {ctes} {rows}"
_REPLACE_SQL: Final = _rewrite_sql('array_replace(array_remove("models", $2), $1, $2)', '$1 = ANY("models")')
_APPEND_SQL: Final = _rewrite_sql('array_append("models", $2)', '$1 = ANY("models") AND NOT ($2 = ANY("models"))')
async def sync_model_allowlists_for_renamed_model(
prisma_client: object,
*,
model_id: str,
old_name: str,
new_name: str,
llm_router: Router | None,
user_api_key_cache: UserApiKeyCache,
) -> None:
if old_name == new_name:
return
executor: Final = raw_executor(prisma_client)
old_name_still_backed: Final = await still_backed(executor, llm_router, old_name, model_id)
touched_rows: Final = await executor.query_raw(
_APPEND_SQL if old_name_still_backed else _REPLACE_SQL, old_name, new_name
)
touched: Final = tuple(_TouchedRow.model_validate(row) for row in touched_rows)
await evict_and_broadcast(
tuple(cache_key for row in touched for cache_key in _CACHE_KEYS_BY_KIND[row.kind](row)),
user_api_key_cache,
)

View file

@ -6106,11 +6106,27 @@ class TestBlockModelResponseSerialization:
class TestAccessGroupModelSync:
"""A rename or delete of a deployment must land in every unified access group that names it."""
"""A rename or delete of a deployment must land in every access group and models allowlist that names it."""
_PS = "litellm.proxy.proxy_server"
_MOD = "litellm.proxy.management_endpoints.model_management_endpoints"
_INVALIDATE = "litellm.proxy.management_helpers.access_group_model_sync.invalidate_access_group_caches"
_EVICT = "litellm.proxy.management_helpers.model_allowlist_rename_sync.evict_and_broadcast"
_ALLOWLIST_TABLES = (
"LiteLLM_TeamTable",
"LiteLLM_VerificationToken",
"LiteLLM_OrganizationTable",
"LiteLLM_ProjectTable",
"LiteLLM_UserTable",
)
_ALLOWLIST_ROWS = [
{"kind": "team", "object_id": "team-1", "team_alias": "alias-1"},
{"kind": "team", "object_id": "team-2", "team_alias": None},
{"kind": "key", "object_id": "hashed-token-1", "team_alias": None},
{"kind": "org", "object_id": "org-1", "team_alias": None},
{"kind": "project", "object_id": "proj-1", "team_alias": None},
{"kind": "user", "object_id": "user-1", "team_alias": None},
]
@staticmethod
def _admin():
@ -6130,7 +6146,10 @@ class TestAccessGroupModelSync:
async def query_raw(sql, *params):
if sql.startswith("SELECT COUNT(*)"):
return [{"deployment_count": deployment_count}]
return [{"access_group_id": "ag-1"}]
if sql.startswith('UPDATE "LiteLLM_AccessGroupTable"'):
return [{"access_group_id": "ag-1"}]
assert sql.startswith("WITH ")
return TestAccessGroupModelSync._ALLOWLIST_ROWS
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
@ -6149,8 +6168,16 @@ class TestAccessGroupModelSync:
if call.args[0].startswith('UPDATE "LiteLLM_AccessGroupTable"')
]
@staticmethod
def _allowlist_updates(mock_prisma):
return [
call
for call in mock_prisma.db.query_raw.await_args_list
if call.args[0].startswith("WITH ") and 'SET "models"' in call.args[0]
]
@contextlib.contextmanager
def _endpoint_env(self, mock_prisma, router):
def _endpoint_env(self, mock_prisma, router, evict=None):
with contextlib.ExitStack() as stack:
for target in (
patch(f"{self._PS}.prisma_client", mock_prisma),
@ -6159,7 +6186,10 @@ class TestAccessGroupModelSync:
patch(f"{self._PS}.premium_user", True),
patch(f"{self._PS}.proxy_logging_obj", MagicMock()),
patch(f"{self._PS}.user_api_key_cache", MagicMock()),
patch(f"{self._MOD}.ModelManagementAuthChecks.can_user_make_model_call", new=AsyncMock(return_value=None)),
patch(self._EVICT, new=evict or AsyncMock()),
patch(
f"{self._MOD}.ModelManagementAuthChecks.can_user_make_model_call", new=AsyncMock(return_value=None)
),
patch(
f"{self._MOD}.clear_cache",
new=AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None)),
@ -6219,7 +6249,9 @@ class TestAccessGroupModelSync:
router.get_model_ids.return_value = ["m-same"]
with self._endpoint_env(mock_prisma, router) as invalidate:
await patch_model(model_id="m-same", patch_data=updateDeployment(blocked=True), user_api_key_dict=self._admin())
await patch_model(
model_id="m-same", patch_data=updateDeployment(blocked=True), user_api_key_dict=self._admin()
)
mock_prisma.db.query_raw.assert_not_awaited()
invalidate.assert_not_awaited()
@ -6280,6 +6312,97 @@ class TestAccessGroupModelSync:
assert update_call.args[1:] == ("gpt-5.6", "gpt-5.6-eu")
invalidate.assert_awaited_once_with(("ag-1",))
@pytest.mark.asyncio
@pytest.mark.parametrize("endpoint", ["patch", "legacy"])
async def test_rename_rewrites_key_team_org_project_and_user_allowlists_and_evicts_their_caches(self, endpoint):
from litellm.proxy.management_endpoints.model_management_endpoints import patch_model, update_model
mock_prisma = self._prisma_with_row("m-rename", "gpt-5.6", deployment_count=0)
router = MagicMock()
router.get_model_ids.return_value = ["m-rename"]
evict = AsyncMock()
with self._endpoint_env(mock_prisma, router, evict=evict):
if endpoint == "patch":
await patch_model(
model_id="m-rename",
patch_data=updateDeployment(model_name="gpt-5.6-eu"),
user_api_key_dict=self._admin(),
)
else:
await update_model(
model_params=updateDeployment(
model_name="gpt-5.6-eu",
litellm_params=updateLiteLLMParams(model="openai/gpt-5.6"),
model_info=ModelInfo(id="m-rename"),
),
user_api_key_dict=self._admin(),
)
(update_call,) = self._allowlist_updates(mock_prisma)
for table in self._ALLOWLIST_TABLES:
assert (
f'UPDATE "{table}" SET "models" = array_replace(array_remove("models", $2), $1, $2) '
'WHERE $1 = ANY("models") RETURNING'
) in update_call.args[0]
assert update_call.args[1:] == ("gpt-5.6", "gpt-5.6-eu")
evict.assert_awaited_once()
assert evict.await_args.args[0] == (
"team_id:team-1",
"team_alias:alias-1",
"team_id:team-2",
"hashed-token-1",
"org_id:org-1",
"org_id:org-1:with_budget",
"project_id:proj-1",
"user-1",
)
@pytest.mark.asyncio
async def test_rename_appends_to_allowlists_when_a_sibling_deployment_keeps_the_old_name(self):
from litellm.proxy.management_endpoints.model_management_endpoints import patch_model
mock_prisma = self._prisma_with_row("m-rename", "gpt-5.6", deployment_count=1)
router = MagicMock()
router.get_model_ids.return_value = ["m-rename"]
with self._endpoint_env(mock_prisma, router):
await patch_model(
model_id="m-rename",
patch_data=updateDeployment(model_name="gpt-5.6-eu"),
user_api_key_dict=self._admin(),
)
(update_call,) = self._allowlist_updates(mock_prisma)
for table in self._ALLOWLIST_TABLES:
assert (
f'UPDATE "{table}" SET "models" = array_append("models", $2) '
'WHERE $1 = ANY("models") AND NOT ($2 = ANY("models")) RETURNING'
) in update_call.args[0]
assert update_call.args[1:] == ("gpt-5.6", "gpt-5.6-eu")
@pytest.mark.asyncio
async def test_unchanged_name_never_touches_allowlists(self):
from litellm.proxy.management_helpers.model_allowlist_rename_sync import (
sync_model_allowlists_for_renamed_model,
)
mock_prisma = self._prisma_with_row("m-rename", "gpt-5.6", deployment_count=0)
evict = AsyncMock()
with patch(self._EVICT, new=evict):
await sync_model_allowlists_for_renamed_model(
prisma_client=mock_prisma,
model_id="m-rename",
old_name="gpt-5.6",
new_name="gpt-5.6",
llm_router=None,
user_api_key_cache=MagicMock(),
)
assert self._allowlist_updates(mock_prisma) == []
evict.assert_not_awaited()
class TestTeamMemberAutoRouterWrites:
@pytest.fixture(autouse=True)