Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/management_helpers/model_allowlist_rename_sync.py: 73%
49 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""
2Keep the `models` allowlists on keys, teams, organizations, projects and users pointing at
3deployment names that still exist.
5Those allowlists store public model names, not ids, so a deployment rename that leaves them
6alone denies the new name while the old entry grants a name nothing serves any more.
7"""
9from collections.abc import Callable
10from dataclasses import dataclass
11from types import MappingProxyType
12from typing import Final
14from pydantic import BaseModel
16from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast
17from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
18from litellm.proxy.management_helpers.access_group_model_sync import raw_executor, still_backed
19from litellm.router import Router
22class _TouchedRow(BaseModel):
23 kind: str
24 object_id: str
25 team_alias: str | None = None
28@dataclass(frozen=True, slots=True)
29class _AllowlistTable:
30 kind: str
31 table: str
32 id_column: str
33 cache_keys: Callable[[_TouchedRow], tuple[str, ...]]
34 alias_column: str | None = None
36 def update_cte(self, set_clause: str, where_clause: str) -> str:
37 alias: Final = f'"{self.alias_column}"' if self.alias_column else "NULL::text"
38 return (
39 f'{self.kind}_rows AS (UPDATE "{self.table}" SET "models" = {set_clause} WHERE {where_clause} '
40 f"RETURNING '{self.kind}' AS kind, \"{self.id_column}\" AS object_id, {alias} AS team_alias)"
41 )
44def _team_cache_keys(row: _TouchedRow) -> tuple[str, ...]:
45 return (f"team_id:{row.object_id}", *((f"team_alias:{row.team_alias}",) if row.team_alias else ()))
48def _key_cache_keys(row: _TouchedRow) -> tuple[str, ...]:
49 return (row.object_id,)
52def _org_cache_keys(row: _TouchedRow) -> tuple[str, ...]:
53 return (f"org_id:{row.object_id}", f"org_id:{row.object_id}:with_budget")
56def _project_cache_keys(row: _TouchedRow) -> tuple[str, ...]:
57 return (f"project_id:{row.object_id}",)
60def _user_cache_keys(row: _TouchedRow) -> tuple[str, ...]:
61 return (row.object_id,)
64_ALLOWLIST_TABLES: Final = (
65 _AllowlistTable("team", "LiteLLM_TeamTable", "team_id", _team_cache_keys, alias_column="team_alias"),
66 _AllowlistTable("key", "LiteLLM_VerificationToken", "token", _key_cache_keys),
67 _AllowlistTable("org", "LiteLLM_OrganizationTable", "organization_id", _org_cache_keys),
68 _AllowlistTable("project", "LiteLLM_ProjectTable", "project_id", _project_cache_keys),
69 _AllowlistTable("user", "LiteLLM_UserTable", "user_id", _user_cache_keys),
70)
72_CACHE_KEYS_BY_KIND: Final = MappingProxyType({table.kind: table.cache_keys for table in _ALLOWLIST_TABLES})
75def _rewrite_sql(set_clause: str, where_clause: str) -> str:
76 """One statement touching every allowlist table, so the rewrite lands everywhere or nowhere."""
77 ctes: Final = ", ".join(table.update_cte(set_clause, where_clause) for table in _ALLOWLIST_TABLES)
78 rows: Final = " UNION ALL ".join(
79 f"SELECT kind, object_id, team_alias FROM {table.kind}_rows" for table in _ALLOWLIST_TABLES
80 )
81 return f"WITH {ctes} {rows}"
84_REPLACE_SQL: Final = _rewrite_sql('array_replace(array_remove("models", $2), $1, $2)', '$1 = ANY("models")')
86_APPEND_SQL: Final = _rewrite_sql('array_append("models", $2)', '$1 = ANY("models") AND NOT ($2 = ANY("models"))')
89async def sync_model_allowlists_for_renamed_model(
90 prisma_client: object,
91 *,
92 model_id: str,
93 old_name: str,
94 new_name: str,
95 llm_router: Router | None,
96 user_api_key_cache: UserApiKeyCache,
97) -> None:
98 if old_name == new_name:
99 return
100 executor: Final = raw_executor(prisma_client)
101 old_name_still_backed: Final = await still_backed(executor, llm_router, old_name, model_id)
102 touched_rows: Final = await executor.query_raw(
103 _APPEND_SQL if old_name_still_backed else _REPLACE_SQL, old_name, new_name
104 )
105 touched: Final = tuple(_TouchedRow.model_validate(row) for row in touched_rows)
106 await evict_and_broadcast(
107 tuple(cache_key for row in touched for cache_key in _CACHE_KEYS_BY_KIND[row.kind](row)),
108 user_api_key_cache,
109 )