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
« 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.
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"""
8from collections.abc import Sequence
9from typing import Final, Protocol
11from pydantic import BaseModel
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
19class _TouchedGroupRow(BaseModel):
20 access_group_id: str
23class _DeploymentCountRow(BaseModel):
24 deployment_count: int
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
31_BACKING_DEPLOYMENTS_SQL: Final = (
32 'SELECT COUNT(*)::int AS deployment_count FROM "LiteLLM_ProxyModelTable" WHERE "model_name" = $1'
33)
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)
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)
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)
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
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
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 )
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)
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 )
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 )
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)