mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(scim): include key permissions in source lifecycle transactions
This commit is contained in:
parent
f4d6091b6c
commit
88f3d22eca
3 changed files with 113 additions and 17 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue