From d3f5cde530d31c93bdda89e3c2113176dfaa931f Mon Sep 17 00:00:00 2001 From: ryan Date: Thu, 17 Sep 2026 23:11:52 +0000 Subject: [PATCH 1/3] fix(proxy): propagate db model renames to key, team, org, project and user model allowlists Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../model_management_endpoints.py | 19 +++ .../access_group_model_sync.py | 16 +-- .../model_allowlist_rename_sync.py | 108 ++++++++++++++++++ .../test_model_management_endpoints.py | 89 ++++++++++++++- 4 files changed, 221 insertions(+), 11 deletions(-) create mode 100644 litellm/proxy/management_helpers/model_allowlist_rename_sync.py diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index bcddb1f7ef0..6208e9eafa6 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -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, @@ -984,6 +985,7 @@ async def patch_model( premium_user, prisma_client, store_model_in_db, + user_api_key_cache, ) try: @@ -1132,6 +1134,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() @@ -2433,6 +2443,7 @@ async def update_model( premium_user, prisma_client, store_model_in_db, + user_api_key_cache, ) try: @@ -2566,6 +2577,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() diff --git a/litellm/proxy/management_helpers/access_group_model_sync.py b/litellm/proxy/management_helpers/access_group_model_sync.py index 7a8dcc2939c..683f2ea79b9 100644 --- a/litellm/proxy/management_helpers/access_group_model_sync.py +++ b/litellm/proxy/management_helpers/access_group_model_sync.py @@ -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) diff --git a/litellm/proxy/management_helpers/model_allowlist_rename_sync.py b/litellm/proxy/management_helpers/model_allowlist_rename_sync.py new file mode 100644 index 00000000000..d0857aae748 --- /dev/null +++ b/litellm/proxy/management_helpers/model_allowlist_rename_sync.py @@ -0,0 +1,108 @@ +""" +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 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 RawExecutor, raw_executor, still_backed +from litellm.router import Router + + +class _TouchedRow(BaseModel): + object_id: str + team_alias: str | None = None + + +@dataclass(frozen=True, slots=True) +class _AllowlistTable: + table: str + returning: str + cache_keys: Callable[[_TouchedRow], tuple[str, ...]] + + def replace_sql(self) -> str: + return ( + f'UPDATE "{self.table}" SET "models" = array_replace(array_remove("models", $2), $1, $2) ' + f'WHERE $1 = ANY("models") RETURNING {self.returning}' + ) + + def append_sql(self) -> str: + return ( + f'UPDATE "{self.table}" SET "models" = array_append("models", $2) ' + f'WHERE $1 = ANY("models") AND NOT ($2 = ANY("models")) RETURNING {self.returning}' + ) + + +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("LiteLLM_TeamTable", '"team_id" AS object_id, "team_alias"', _team_cache_keys), + _AllowlistTable("LiteLLM_VerificationToken", '"token" AS object_id', _key_cache_keys), + _AllowlistTable("LiteLLM_OrganizationTable", '"organization_id" AS object_id', _org_cache_keys), + _AllowlistTable("LiteLLM_ProjectTable", '"project_id" AS object_id', _project_cache_keys), + _AllowlistTable("LiteLLM_UserTable", '"user_id" AS object_id', _user_cache_keys), +) + + +async def _rewrite_allowlist( + executor: RawExecutor, + allowlist: _AllowlistTable, + sql: str, + old_name: str, + new_name: str, + user_api_key_cache: UserApiKeyCache, +) -> None: + touched_rows: Final = await executor.query_raw(sql, old_name, new_name) + await evict_and_broadcast( + tuple(cache_key for row in touched_rows for cache_key in allowlist.cache_keys(_TouchedRow.model_validate(row))), + user_api_key_cache, + ) + + +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: + executor: Final = raw_executor(prisma_client) + old_name_still_backed: Final = await still_backed(executor, llm_router, old_name, model_id) + for allowlist in _ALLOWLIST_TABLES: + await _rewrite_allowlist( + executor, + allowlist, + allowlist.append_sql() if old_name_still_backed else allowlist.replace_sql(), + old_name, + new_name, + user_api_key_cache, + ) diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index e46b4fee61c..d164fa28c10 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -6058,11 +6058,19 @@ 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_ROWS = { + "LiteLLM_TeamTable": [{"object_id": "team-1", "team_alias": "alias-1"}, {"object_id": "team-2", "team_alias": None}], + "LiteLLM_VerificationToken": [{"object_id": "hashed-token-1"}], + "LiteLLM_OrganizationTable": [{"object_id": "org-1"}], + "LiteLLM_ProjectTable": [{"object_id": "proj-1"}], + "LiteLLM_UserTable": [{"object_id": "user-1"}], + } @staticmethod def _admin(): @@ -6082,7 +6090,9 @@ 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"}] + return TestAccessGroupModelSync._ALLOWLIST_ROWS[sql.split('"')[1]] mock_prisma = MagicMock() mock_prisma.db = MagicMock() @@ -6101,8 +6111,16 @@ class TestAccessGroupModelSync: if call.args[0].startswith('UPDATE "LiteLLM_AccessGroupTable"') ] + @staticmethod + def _allowlist_updates(mock_prisma): + return { + call.args[0].split('"')[1]: call + for call in mock_prisma.db.query_raw.await_args_list + if call.args[0].startswith('UPDATE "') 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), @@ -6111,6 +6129,7 @@ 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(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", @@ -6232,6 +6251,70 @@ 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(), + ) + + updates = self._allowlist_updates(mock_prisma) + assert set(updates) == set(self._ALLOWLIST_ROWS) + for update_call in updates.values(): + assert 'SET "models" = array_replace(array_remove("models", $2), $1, $2)' in update_call.args[0] + assert 'WHERE $1 = ANY("models")' in update_call.args[0] + assert update_call.args[1:] == ("gpt-5.6", "gpt-5.6-eu") + evicted = [call.args[0] for call in evict.await_args_list] + assert evicted == [ + ("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(), + ) + + updates = self._allowlist_updates(mock_prisma) + assert set(updates) == set(self._ALLOWLIST_ROWS) + for update_call in updates.values(): + assert 'SET "models" = array_append("models", $2)' in update_call.args[0] + assert 'WHERE $1 = ANY("models") AND NOT ($2 = ANY("models"))' in update_call.args[0] + assert update_call.args[1:] == ("gpt-5.6", "gpt-5.6-eu") + class TestTeamMemberAutoRouterWrites: @pytest.fixture(autouse=True) From de0047c802d906251b6fde03b4a1a7ebf1c3cd8e Mon Sep 17 00:00:00 2001 From: ryan Date: Thu, 17 Sep 2026 23:24:26 +0000 Subject: [PATCH 2/3] fix(proxy): skip allowlist rewrite when the model name is unchanged Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../model_allowlist_rename_sync.py | 2 ++ .../test_model_management_endpoints.py | 22 +++++++++++++++++++ 2 files changed, 24 insertions(+) diff --git a/litellm/proxy/management_helpers/model_allowlist_rename_sync.py b/litellm/proxy/management_helpers/model_allowlist_rename_sync.py index d0857aae748..1e92c383f01 100644 --- a/litellm/proxy/management_helpers/model_allowlist_rename_sync.py +++ b/litellm/proxy/management_helpers/model_allowlist_rename_sync.py @@ -95,6 +95,8 @@ async def sync_model_allowlists_for_renamed_model( 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) for allowlist in _ALLOWLIST_TABLES: diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index d164fa28c10..0f273ebce84 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -6315,6 +6315,28 @@ class TestAccessGroupModelSync: assert 'WHERE $1 = ANY("models") AND NOT ($2 = ANY("models"))' 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) From 1d50d1ad3b9caf69f05d7c2209ec2862da49171e Mon Sep 17 00:00:00 2001 From: ryan Date: Thu, 17 Sep 2026 23:40:59 +0000 Subject: [PATCH 3/3] fix(proxy): rewrite every model allowlist in one statement on rename Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../model_allowlist_rename_sync.py | 73 ++++++++------- .../test_model_management_endpoints.py | 88 +++++++++++-------- 2 files changed, 89 insertions(+), 72 deletions(-) diff --git a/litellm/proxy/management_helpers/model_allowlist_rename_sync.py b/litellm/proxy/management_helpers/model_allowlist_rename_sync.py index 1e92c383f01..f93312f7a37 100644 --- a/litellm/proxy/management_helpers/model_allowlist_rename_sync.py +++ b/litellm/proxy/management_helpers/model_allowlist_rename_sync.py @@ -8,37 +8,36 @@ alone denies the new name while the old entry grants a name nothing serves any m 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 RawExecutor, raw_executor, still_backed +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 - returning: str + id_column: str cache_keys: Callable[[_TouchedRow], tuple[str, ...]] + alias_column: str | None = None - def replace_sql(self) -> str: + 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'UPDATE "{self.table}" SET "models" = array_replace(array_remove("models", $2), $1, $2) ' - f'WHERE $1 = ANY("models") RETURNING {self.returning}' - ) - - def append_sql(self) -> str: - return ( - f'UPDATE "{self.table}" SET "models" = array_append("models", $2) ' - f'WHERE $1 = ANY("models") AND NOT ($2 = ANY("models")) RETURNING {self.returning}' + 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)" ) @@ -63,27 +62,28 @@ def _user_cache_keys(row: _TouchedRow) -> tuple[str, ...]: _ALLOWLIST_TABLES: Final = ( - _AllowlistTable("LiteLLM_TeamTable", '"team_id" AS object_id, "team_alias"', _team_cache_keys), - _AllowlistTable("LiteLLM_VerificationToken", '"token" AS object_id', _key_cache_keys), - _AllowlistTable("LiteLLM_OrganizationTable", '"organization_id" AS object_id', _org_cache_keys), - _AllowlistTable("LiteLLM_ProjectTable", '"project_id" AS object_id', _project_cache_keys), - _AllowlistTable("LiteLLM_UserTable", '"user_id" AS object_id', _user_cache_keys), + _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}) -async def _rewrite_allowlist( - executor: RawExecutor, - allowlist: _AllowlistTable, - sql: str, - old_name: str, - new_name: str, - user_api_key_cache: UserApiKeyCache, -) -> None: - touched_rows: Final = await executor.query_raw(sql, old_name, new_name) - await evict_and_broadcast( - tuple(cache_key for row in touched_rows for cache_key in allowlist.cache_keys(_TouchedRow.model_validate(row))), - user_api_key_cache, + +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( @@ -99,12 +99,11 @@ async def sync_model_allowlists_for_renamed_model( return executor: Final = raw_executor(prisma_client) old_name_still_backed: Final = await still_backed(executor, llm_router, old_name, model_id) - for allowlist in _ALLOWLIST_TABLES: - await _rewrite_allowlist( - executor, - allowlist, - allowlist.append_sql() if old_name_still_backed else allowlist.replace_sql(), - old_name, - new_name, - user_api_key_cache, - ) + 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, + ) diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index 0f273ebce84..20e51e7c906 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -6064,13 +6064,21 @@ class TestAccessGroupModelSync: _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_ROWS = { - "LiteLLM_TeamTable": [{"object_id": "team-1", "team_alias": "alias-1"}, {"object_id": "team-2", "team_alias": None}], - "LiteLLM_VerificationToken": [{"object_id": "hashed-token-1"}], - "LiteLLM_OrganizationTable": [{"object_id": "org-1"}], - "LiteLLM_ProjectTable": [{"object_id": "proj-1"}], - "LiteLLM_UserTable": [{"object_id": "user-1"}], - } + _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(): @@ -6092,7 +6100,8 @@ class TestAccessGroupModelSync: return [{"deployment_count": deployment_count}] if sql.startswith('UPDATE "LiteLLM_AccessGroupTable"'): return [{"access_group_id": "ag-1"}] - return TestAccessGroupModelSync._ALLOWLIST_ROWS[sql.split('"')[1]] + assert sql.startswith("WITH ") + return TestAccessGroupModelSync._ALLOWLIST_ROWS mock_prisma = MagicMock() mock_prisma.db = MagicMock() @@ -6113,11 +6122,11 @@ class TestAccessGroupModelSync: @staticmethod def _allowlist_updates(mock_prisma): - return { - call.args[0].split('"')[1]: call + return [ + call for call in mock_prisma.db.query_raw.await_args_list - if call.args[0].startswith('UPDATE "') and 'SET "models"' in call.args[0] - } + if call.args[0].startswith("WITH ") and 'SET "models"' in call.args[0] + ] @contextlib.contextmanager def _endpoint_env(self, mock_prisma, router, evict=None): @@ -6130,7 +6139,9 @@ class TestAccessGroupModelSync: patch(f"{self._PS}.proxy_logging_obj", MagicMock()), patch(f"{self._PS}.user_api_key_cache", MagicMock()), 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}.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)), @@ -6190,7 +6201,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() @@ -6278,20 +6291,24 @@ class TestAccessGroupModelSync: user_api_key_dict=self._admin(), ) - updates = self._allowlist_updates(mock_prisma) - assert set(updates) == set(self._ALLOWLIST_ROWS) - for update_call in updates.values(): - assert 'SET "models" = array_replace(array_remove("models", $2), $1, $2)' in update_call.args[0] - assert 'WHERE $1 = ANY("models")' in update_call.args[0] - assert update_call.args[1:] == ("gpt-5.6", "gpt-5.6-eu") - evicted = [call.args[0] for call in evict.await_args_list] - assert evicted == [ - ("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",), - ] + (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): @@ -6308,12 +6325,13 @@ class TestAccessGroupModelSync: user_api_key_dict=self._admin(), ) - updates = self._allowlist_updates(mock_prisma) - assert set(updates) == set(self._ALLOWLIST_ROWS) - for update_call in updates.values(): - assert 'SET "models" = array_append("models", $2)' in update_call.args[0] - assert 'WHERE $1 = ANY("models") AND NOT ($2 = ANY("models"))' in update_call.args[0] - assert update_call.args[1:] == ("gpt-5.6", "gpt-5.6-eu") + (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): @@ -6334,7 +6352,7 @@ class TestAccessGroupModelSync: user_api_key_cache=MagicMock(), ) - assert self._allowlist_updates(mock_prisma) == {} + assert self._allowlist_updates(mock_prisma) == [] evict.assert_not_awaited()