From 88f3d22eca5d0f9fa259e6eea664917b54f95716 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 22:29:11 -0700 Subject: [PATCH] fix(scim): include key permissions in source lifecycle transactions --- .../key_management_endpoints.py | 45 ++++++++---- .../object_permission_utils.py | 16 ++++- .../test_key_management_endpoints.py | 69 +++++++++++++++++++ 3 files changed, 113 insertions(+), 17 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index aded71edd78..bc4ee68e011 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -2437,16 +2437,22 @@ async def _update_key_row_with_soft_budget( async with prisma_client.tx() as tx: if "allowed_routes" in data.model_fields_set: await _lock_and_validate_source_key_change(tx, hashed_token, data.allowed_routes) + permission_values: Final = await _handle_update_object_permission( + data_json=dict(non_default_values), + existing_key_row=existing_key_row, + prisma_client=prisma_client, + tx=tx, + ) update_values: Final = ( await _apply_soft_budget_update( data=data, - non_default_values=non_default_values, + non_default_values=permission_values, db=tx, existing_key_row=existing_key_row, changed_by=changed_by, ) if "soft_budget" in data.model_fields_set - else non_default_values + else permission_values ) include_object_permission: Final[prisma.types.LiteLLM_VerificationTokenInclude] = {"object_permission": True} updated_row: Final = await tx.litellm_verificationtoken.update( @@ -2573,6 +2579,8 @@ async def _handle_update_object_permission( data_json: dict, existing_key_row: LiteLLM_VerificationToken, prisma_client: PrismaClient, + *, + tx: "Prisma | None" = None, ) -> dict: """Persist the requested object permission row and swap it for its id, only after the key policy allowed the write.""" if "object_permission" not in data_json: @@ -2582,6 +2590,7 @@ async def _handle_update_object_permission( data_json=data_json, existing_object_permission_id=existing_key_row.object_permission_id, prisma_client=prisma_client, + tx=tx, ) # Add the object_permission_id to data_json if one was created/updated @@ -3544,10 +3553,15 @@ async def update_key_fn( if prisma_client is None: raise Exception("Not connected to DB!") - update_values: Final = await _handle_update_object_permission( - data_json=non_default_values, - existing_key_row=existing_key_row, - prisma_client=prisma_client, + uses_transaction: Final = bool(data.model_fields_set.intersection(("soft_budget", "allowed_routes"))) + update_values: Final = ( + non_default_values + if uses_transaction + else await _handle_update_object_permission( + data_json=non_default_values, + existing_key_row=existing_key_row, + prisma_client=prisma_client, + ) ) changed_by: Final = user_api_key_dict.user_id or litellm_proxy_admin_name response: Final = ( @@ -3559,7 +3573,7 @@ async def update_key_fn( existing_key_row=existing_key_row, changed_by=changed_by, ) - if {"soft_budget", "allowed_routes"}.intersection(data.model_fields_set) + if uses_transaction else await prisma_client.update_data(token=key, data=MappingProxyType({**update_values, "token": key})) ) @@ -5678,14 +5692,6 @@ async def _execute_virtual_key_regeneration( request=data if data is not None else RegenerateKeyRequest(), ), ) - update_values: Final = await _handle_update_object_permission( - data_json=non_default_values, - existing_key_row=key_in_db, - prisma_client=prisma_client, - ) - update_data.update(update_values) - jsonified_update_data: Final[Mapping[str, object]] = prisma_client.jsonify_object(data=update_data) - # Snapshot before the token update: the FK cascade rewrites mapping rows to the new hash, # but their cached jwt_key_mapping entries still point at the old token (LIT-5379). jwt_mapping_cache_keys: Final = await get_jwt_key_mapping_cache_keys_for_token( @@ -5695,6 +5701,15 @@ async def _execute_virtual_key_regeneration( async with prisma_client.tx() as tx: await _lock_and_validate_source_key_change(tx, hashed_api_key) + update_values: Final = await _handle_update_object_permission( + data_json=non_default_values, + existing_key_row=key_in_db, + prisma_client=prisma_client, + tx=tx, + ) + jsonified_update_data: Final[Mapping[str, object]] = prisma_client.jsonify_object( + data={**update_data, **update_values} + ) await _persist_deleted_verification_tokens( keys=[key_in_db], prisma_client=prisma_client, diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index 569f2ebd278..68ac96a609d 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -25,6 +25,7 @@ from litellm.repositories.prisma_protocols import DatabaseClient, TableActions from litellm.repositories.table_repositories import MCPServerRepository if TYPE_CHECKING: + from prisma import Prisma from prisma import models as prisma_models from litellm.proxy._types import ( @@ -84,6 +85,8 @@ async def prepare_object_permission_upsert( new_object_permission: Mapping[str, object], existing_object_permission_id: str | None, prisma_client: PrismaClient, + *, + tx: "Prisma | None" = None, ) -> ObjectPermissionUpsert: """ Read-and-merge half of an object permission upsert; performs no writes. @@ -101,7 +104,10 @@ async def prepare_object_permission_upsert( update cannot leave permission changes live. """ object_permission_id: Final = existing_object_permission_id or str(uuid.uuid4()) - existing_object_permission: Final = await ObjectPermissionRepository(prisma_client).table.find_unique( + permission_table: Final = ( + tx.litellm_objectpermissiontable if tx is not None else ObjectPermissionRepository(prisma_client).table + ) + existing_object_permission: Final = await permission_table.find_unique( where={"object_permission_id": object_permission_id}, ) existing_fields: Final[dict[str, object]] = ( @@ -134,6 +140,8 @@ async def handle_update_object_permission_common( data_json: dict, existing_object_permission_id: str | None, prisma_client: PrismaClient | None, + *, + tx: "Prisma | None" = None, ) -> str | None: """ Common logic for handling object permission updates across organizations, teams, and keys. @@ -170,8 +178,12 @@ async def handle_update_object_permission_common( new_object_permission=new_object_permission if isinstance(new_object_permission, dict) else {}, existing_object_permission_id=existing_object_permission_id, prisma_client=prisma_client, + tx=tx, ) - created_object_permission_row: Final = await ObjectPermissionRepository(prisma_client).table.upsert( + permission_table: Final = ( + tx.litellm_objectpermissiontable if tx is not None else ObjectPermissionRepository(prisma_client).table + ) + created_object_permission_row: Final = await permission_table.upsert( where={"object_permission_id": upsert.object_permission_id}, data={ "create": upsert.record, diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index d011cd8cbee..2e0ce748bcd 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -21258,3 +21258,72 @@ async def test_rotation_grace_period_is_written_to_the_selected_database(transac assert saved["data"]["update"]["active_token_id"] == "new-token" assert before + timedelta(hours=1) <= saved["data"]["create"]["revoke_at"] <= datetime.now(timezone.utc) + timedelta(hours=1) unused.litellm_deprecatedverificationtoken.upsert.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["update", "regenerate"]) +async def test_source_binding_race_denial_preserves_object_permissions(monkeypatch, operation): + from types import SimpleNamespace + from litellm.proxy import proxy_server + from litellm.proxy.management_endpoints.key_management_endpoints import update_key_fn, _execute_virtual_key_regeneration + + key = _make_regenerate_existing_key().model_copy(update={"allowed_routes": ["/scim/*"], "object_permission_id": "permission-before"}) + if operation == "update": + _wire_update_key_fn(monkeypatch, key) + client = proxy_server.prisma_client + else: + client = _make_regenerate_mock_prisma() + client.writer_db.litellm_scimsource.find_unique = AsyncMock(return_value=None) + events = [] + _record_object_permission_writes(client, events) + tx = AsyncMock() + client.tx = MagicMock() + client.tx.return_value.__aenter__.return_value = tx + tx.litellm_verificationtoken.find_unique.return_value = key + tx.litellm_scimsource.find_unique.return_value = SimpleNamespace(source_id="concurrent-source") + permission = LiteLLM_ObjectPermissionBase(vector_stores=["vs-after"]) + if operation == "update": + request = MagicMock() + request.query_params = {} + call = update_key_fn(request=request, data=UpdateKeyRequest(key=key.token, allowed_routes=["/*"], object_permission=permission), user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None) + else: + call = _execute_virtual_key_regeneration(prisma_client=client, key_in_db=key, hashed_api_key=key.token, key="sk-original", data=RegenerateKeyRequest(object_permission=permission), user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None, user_api_key_cache=MagicMock(), proxy_logging_obj=MagicMock()) + with _patch_regenerate_side_effects(): + with pytest.raises((HTTPException, ProxyException)) as denied: + await call + assert str(getattr(denied.value, "code", getattr(denied.value, "status_code", None))) == "409" + assert events == [] + client.db.litellm_objectpermissiontable.upsert.assert_not_awaited() + tx.litellm_objectpermissiontable.upsert.assert_not_awaited() + tx.litellm_verificationtoken.update.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["update", "regenerate"]) +async def test_permission_update_joins_key_write_transaction(operation): + from litellm.proxy.management_endpoints.key_management_endpoints import _execute_virtual_key_regeneration, _update_key_row_with_soft_budget + + client = _make_regenerate_mock_prisma() + client.writer_db.litellm_scimsource.find_unique = AsyncMock(return_value=None) + key = _make_regenerate_existing_key().model_copy(update={"object_permission_id": "permission-before"}) + tx = AsyncMock() + client.tx.return_value.__aenter__.return_value = tx + tx.litellm_verificationtoken.find_unique.return_value = key + tx.litellm_scimsource.find_unique.return_value = None + tx.litellm_objectpermissiontable.find_unique.return_value = None + tx.litellm_objectpermissiontable.upsert.return_value = MagicMock(object_permission_id="permission-before") + tx.litellm_verificationtoken.update.side_effect = RuntimeError("token write failed") + permission = LiteLLM_ObjectPermissionBase(vector_stores=["vs-after"]) + if operation == "update": + call = _update_key_row_with_soft_budget(client, key.token, UpdateKeyRequest(key=key.token, allowed_routes=[], object_permission=permission), {"allowed_routes": [], "object_permission": permission.model_dump()}, key, "admin") + else: + call = _execute_virtual_key_regeneration(prisma_client=client, key_in_db=key, hashed_api_key=key.token, key="sk-original", data=RegenerateKeyRequest(object_permission=permission), user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None, user_api_key_cache=MagicMock(), proxy_logging_obj=MagicMock()) + with _patch_regenerate_side_effects(): + with pytest.raises(RuntimeError, match="token write failed"): + await call + client.db.litellm_objectpermissiontable.upsert.assert_not_awaited() + saved = tx.litellm_objectpermissiontable.upsert.await_args.kwargs + assert saved["where"] == {"object_permission_id": "permission-before"} + assert saved["data"]["update"]["vector_stores"] == ["vs-after"] + assert tx.litellm_verificationtoken.update.await_args.kwargs["data"]["object_permission_id"] == "permission-before" + assert client.tx.return_value.__aexit__.await_args.args[0] is RuntimeError