Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/management_helpers/access_group_model_sync.py: 63%

48 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1""" 

2Keep `litellm_accessgrouptable.access_model_names` pointing at deployment names that still exist. 

3 

4Unified access groups store model names, not ids, so a deployment rename or delete that leaves 

5the arrays alone strands every group on a name nothing serves any more. 

6""" 

7 

8from collections.abc import Sequence 

9from typing import Final, Protocol 

10 

11from pydantic import BaseModel 

12 

13from litellm.proxy.db.routing_prisma_wrapper import writer_wrapper 

14from litellm.proxy.management_helpers.access_group_team_sync import invalidate_access_group_caches 

15from litellm.repositories.table_repositories import AccessGroupRepository 

16from litellm.router import Router 

17 

18 

19class _TouchedGroupRow(BaseModel): 

20 access_group_id: str 

21 

22 

23class _DeploymentCountRow(BaseModel): 

24 deployment_count: int 

25 

26 

27class RawExecutor(Protocol): 

28 async def query_raw(self, query: str, *args: str) -> Sequence[object]: ... 28 ↛ exitline 28 didn't return from function 'query_raw' because

29 

30 

31_BACKING_DEPLOYMENTS_SQL: Final = ( 

32 'SELECT COUNT(*)::int AS deployment_count FROM "LiteLLM_ProxyModelTable" WHERE "model_name" = $1' 

33) 

34 

35_REPLACE_MODEL_NAME_SQL: Final = ( 

36 'UPDATE "LiteLLM_AccessGroupTable" ' 

37 'SET "access_model_names" = array_replace(array_remove("access_model_names", $2), $1, $2) ' 

38 'WHERE $1 = ANY("access_model_names") ' 

39 'RETURNING "access_group_id"' 

40) 

41 

42_APPEND_MODEL_NAME_SQL: Final = ( 

43 'UPDATE "LiteLLM_AccessGroupTable" ' 

44 'SET "access_model_names" = array_append("access_model_names", $2) ' 

45 'WHERE $1 = ANY("access_model_names") AND NOT ($2 = ANY("access_model_names")) ' 

46 'RETURNING "access_group_id"' 

47) 

48 

49_REMOVE_MODEL_NAME_SQL: Final = ( 

50 'UPDATE "LiteLLM_AccessGroupTable" ' 

51 'SET "access_model_names" = array_remove("access_model_names", $1) ' 

52 'WHERE $1 = ANY("access_model_names") ' 

53 'RETURNING "access_group_id"' 

54) 

55 

56 

57def raw_executor(prisma_client: object) -> RawExecutor: 

58 db: Final = AccessGroupRepository(prisma_client).prisma_client.db # pyright: ignore[reportAny] # untyped Prisma client 

59 return writer_wrapper(db) # pyright: ignore[reportAny, reportReturnType] # untyped Prisma client behind the pin 

60 

61 

62def _config_sourced_sibling(llm_router: Router, deployment_id: str, model_id: str) -> bool: 

63 if deployment_id == model_id: 

64 return False 

65 deployment: Final = llm_router.get_deployment(model_id=deployment_id) 

66 return deployment is not None and not deployment.model_info.db_model 

67 

68 

69def _served_by_a_config_deployment(llm_router: Router | None, model_name: str, model_id: str) -> bool: 

70 if llm_router is None: 70 ↛ 71line 70 didn't jump to line 71 because the condition on line 70 was never true

71 return False 

72 return any( 

73 _config_sourced_sibling(llm_router, deployment_id, model_id) 

74 for deployment_id in llm_router.get_model_ids(model_name=model_name) 

75 ) 

76 

77 

78async def still_backed(executor: RawExecutor, llm_router: Router | None, model_name: str, model_id: str) -> bool: 

79 if _served_by_a_config_deployment(llm_router, model_name, model_id): 79 ↛ 80line 79 didn't jump to line 80 because the condition on line 79 was never true

80 return True 

81 count_rows: Final = await executor.query_raw(_BACKING_DEPLOYMENTS_SQL, model_name) 

82 return any(_DeploymentCountRow.model_validate(row).deployment_count > 0 for row in count_rows) 

83 

84 

85async def _rewrite_groups(executor: RawExecutor, sql: str, *names: str) -> None: 

86 touched_rows: Final = await executor.query_raw(sql, *names) 

87 await invalidate_access_group_caches( 

88 tuple(_TouchedGroupRow.model_validate(row).access_group_id for row in touched_rows) 

89 ) 

90 

91 

92async def sync_access_groups_for_renamed_model( 

93 prisma_client: object, 

94 *, 

95 model_id: str, 

96 old_name: str, 

97 new_name: str, 

98 llm_router: Router | None, 

99) -> None: 

100 if old_name == new_name: 

101 return 

102 executor: Final = raw_executor(prisma_client) 

103 old_name_still_backed: Final = await still_backed(executor, llm_router, old_name, model_id) 

104 await _rewrite_groups( 

105 executor, _APPEND_MODEL_NAME_SQL if old_name_still_backed else _REPLACE_MODEL_NAME_SQL, old_name, new_name 

106 ) 

107 

108 

109async def sync_access_groups_for_deleted_model( 

110 prisma_client: object, 

111 *, 

112 model_id: str, 

113 model_name: str, 

114 llm_router: Router | None, 

115) -> None: 

116 executor: Final = raw_executor(prisma_client) 

117 if await still_backed(executor, llm_router, model_name, model_id): 117 ↛ 119line 117 didn't jump to line 119 because the condition on line 117 was always true

118 return 

119 await _rewrite_groups(executor, _REMOVE_MODEL_NAME_SQL, model_name)