mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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>
This commit is contained in:
parent
07b5051c0d
commit
d3f5cde530
4 changed files with 221 additions and 11 deletions
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
108
litellm/proxy/management_helpers/model_allowlist_rename_sync.py
Normal file
108
litellm/proxy/management_helpers/model_allowlist_rename_sync.py
Normal file
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue