fix(scim): include key permissions in source lifecycle transactions
Some checks are pending
LiteLLM Rust / rust-lint (push) Waiting to run
LiteLLM Rust / rust-test (push) Waiting to run
LiteLLM Rust / rust-wheel (push) Waiting to run

This commit is contained in:
Joshua Valluru 2026-09-28 22:29:11 -07:00
parent f4d6091b6c
commit 88f3d22eca
3 changed files with 113 additions and 17 deletions

View file

@ -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,

View file

@ -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,

View file

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