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

1""" 

2Keep the `models` allowlists on keys, teams, organizations, projects and users pointing at 

3deployment names that still exist. 

4 

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""" 

8 

9from collections.abc import Callable 

10from dataclasses import dataclass 

11from types import MappingProxyType 

12from typing import Final 

13 

14from pydantic import BaseModel 

15 

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 

20 

21 

22class _TouchedRow(BaseModel): 

23 kind: str 

24 object_id: str 

25 team_alias: str | None = None 

26 

27 

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 

35 

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 ) 

42 

43 

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 ())) 

46 

47 

48def _key_cache_keys(row: _TouchedRow) -> tuple[str, ...]: 

49 return (row.object_id,) 

50 

51 

52def _org_cache_keys(row: _TouchedRow) -> tuple[str, ...]: 

53 return (f"org_id:{row.object_id}", f"org_id:{row.object_id}:with_budget") 

54 

55 

56def _project_cache_keys(row: _TouchedRow) -> tuple[str, ...]: 

57 return (f"project_id:{row.object_id}",) 

58 

59 

60def _user_cache_keys(row: _TouchedRow) -> tuple[str, ...]: 

61 return (row.object_id,) 

62 

63 

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) 

71 

72_CACHE_KEYS_BY_KIND: Final = MappingProxyType({table.kind: table.cache_keys for table in _ALLOWLIST_TABLES}) 

73 

74 

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}" 

82 

83 

84_REPLACE_SQL: Final = _rewrite_sql('array_replace(array_remove("models", $2), $1, $2)', '$1 = ANY("models")') 

85 

86_APPEND_SQL: Final = _rewrite_sql('array_append("models", $2)', '$1 = ANY("models") AND NOT ($2 = ANY("models"))') 

87 

88 

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 )