From d362a91d2b31c0a6f3769fdd451b0bc4cb1c6b76 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 21:52:40 -0700 Subject: [PATCH] fix(scim): serialize source binding with token mutations --- .../key_management_endpoints.py | 91 +++++++++++++------ .../scim/source_endpoints.py | 7 +- .../scim/test_source_endpoints.py | 16 ++++ .../test_key_management_endpoints.py | 85 +++++++++++++++++ 4 files changed, 168 insertions(+), 31 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index a8124ac6a12..aded71edd78 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -2434,14 +2434,19 @@ async def _update_key_row_with_soft_budget( ) -> _KeyUpdateResult: hashed_token: Final = _hash_token_if_needed(key) key_where: Final[_KeyRowWhere] = {"token": hashed_token} - tx: _KeyUpdateTx async with prisma_client.tx() as tx: - update_values: Final = await _apply_soft_budget_update( - data=data, - non_default_values=non_default_values, - db=tx, - existing_key_row=existing_key_row, - changed_by=changed_by, + if "allowed_routes" in data.model_fields_set: + await _lock_and_validate_source_key_change(tx, hashed_token, data.allowed_routes) + update_values: Final = ( + await _apply_soft_budget_update( + data=data, + non_default_values=non_default_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 ) include_object_permission: Final[prisma.types.LiteLLM_VerificationTokenInclude] = {"object_permission": True} updated_row: Final = await tx.litellm_verificationtoken.update( @@ -3498,6 +3503,9 @@ async def update_key_fn( await _enforce_custom_key_update_policy(hook=_custom_key_update_hook(proxy_server), data=data) + if "allowed_routes" in data.model_fields_set and tuple(data.allowed_routes or ()) != ("/scim/*",): + await _reject_source_bound_key_change(prisma_client, existing_key_row) + # Enforce upperbound key params on update (don't fill defaults) _enforce_upperbound_key_params(data, fill_defaults=False) non_default_values: Final = await prepare_key_update_data( @@ -3551,7 +3559,7 @@ async def update_key_fn( existing_key_row=existing_key_row, changed_by=changed_by, ) - if "soft_budget" in data.model_fields_set + if {"soft_budget", "allowed_routes"}.intersection(data.model_fields_set) else await prisma_client.update_data(token=key, data=MappingProxyType({**update_values, "token": key})) ) @@ -5484,6 +5492,7 @@ async def _insert_deprecated_key( old_token_hash: str, new_token_hash: str, grace_period: str | None, + tx: "Prisma | None" = None, ) -> None: """ Insert old key into deprecated table so it remains valid during grace period. @@ -5514,7 +5523,12 @@ async def _insert_deprecated_key( try: revoke_at: Final = datetime.now(timezone.utc) + timedelta(seconds=grace_seconds) - await _deprecated_verification_token_table(prisma_client).upsert( + table: Final = ( + tx.litellm_deprecatedverificationtoken + if tx is not None + else _deprecated_verification_token_table(prisma_client) + ) + await table.upsert( where={"token": old_token_hash}, data={ "create": { @@ -5540,6 +5554,24 @@ async def _insert_deprecated_key( ) +async def _lock_and_validate_source_key_change( + tx: "Prisma", token: str, allowed_routes: Sequence[str] | None = None +) -> None: + from litellm.proxy.management_endpoints.scim.source_endpoints import lock_provisioning_token + + await lock_provisioning_token(tx, token) + key: Final = await tx.litellm_verificationtoken.find_unique(where={"token": token}) + if key is None: + raise HTTPException(409, "The key changed during the request; retry with the current key") + if tuple(allowed_routes or ()) == ("/scim/*",): + return + source: Final = await tx.litellm_scimsource.find_unique(where={"key_hash": token}) + if source is not None: + raise HTTPException( + 409, "A provisioning source token cannot be regenerated or have its SCIM restriction removed" + ) + + async def _reject_source_bound_key_change(prisma_client: PrismaClient, key: LiteLLM_VerificationToken) -> None: if tuple(key.allowed_routes or ()) != ("/scim/*",): return @@ -5661,27 +5693,26 @@ async def _execute_virtual_key_regeneration( prisma_client=prisma_client, ) - await _persist_deleted_verification_tokens( - keys=[key_in_db], - prisma_client=prisma_client, - user_api_key_dict=user_api_key_dict, - litellm_changed_by=litellm_changed_by, - ) - - # If grace period set, insert deprecated key so old key remains valid - await _insert_deprecated_key( - prisma_client=prisma_client, - old_token_hash=hashed_api_key, - new_token_hash=new_token_hash, - grace_period=data.grace_period if data else None, - ) - - updated_token: Final[LiteLLM_VerificationToken | None] = await _prisma_table( - VerificationTokenRepository(prisma_client) - ).update( - where={"token": hashed_api_key}, - data=with_settings_updated_at(jsonified_update_data), - ) + async with prisma_client.tx() as tx: + await _lock_and_validate_source_key_change(tx, hashed_api_key) + await _persist_deleted_verification_tokens( + keys=[key_in_db], + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + litellm_changed_by=litellm_changed_by, + tx=tx, + ) + await _insert_deprecated_key( + prisma_client=prisma_client, + old_token_hash=hashed_api_key, + new_token_hash=new_token_hash, + grace_period=data.grace_period if data else None, + tx=tx, + ) + updated_token: Final = await tx.litellm_verificationtoken.update( + where={"token": hashed_api_key}, + data=with_settings_updated_at(jsonified_update_data), + ) updated_token_dict: Final[dict[str, object]] = dict(updated_token) if updated_token is not None else {} updated_token_dict["key"] = new_token updated_token_dict["token_id"] = updated_token_dict.pop("token") diff --git a/litellm/proxy/management_endpoints/scim/source_endpoints.py b/litellm/proxy/management_endpoints/scim/source_endpoints.py index fb7136f0a12..0bd78ccb195 100644 --- a/litellm/proxy/management_endpoints/scim/source_endpoints.py +++ b/litellm/proxy/management_endpoints/scim/source_endpoints.py @@ -4,7 +4,7 @@ from typing import Annotated, Final from uuid import uuid4 from fastapi import APIRouter, Depends, HTTPException -from prisma import Json +from prisma import Json, Prisma from prisma.types import ( LiteLLM_SCIMSourceCreateInput, LiteLLM_SCIMSourceOrderByInput, @@ -62,6 +62,10 @@ async def list_sources(auth: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth return tuple(source_response(source) for source in sources) +async def lock_provisioning_token(tx: Prisma, token_hash: str) -> None: + await tx.execute_raw('SELECT 1 FROM "LiteLLM_VerificationToken" WHERE token = $1 FOR UPDATE', token_hash) + + @router.post("", response_model=SCIMSourceResponse, status_code=201) async def create_source( data: SCIMSourceCreate, auth: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)] @@ -71,6 +75,7 @@ async def create_source( token_hash: Final = hash_token(data.provisioning_token.get_secret_value()) source_filter: Final[LiteLLM_SCIMSourceWhereUniqueInput] = {"key_hash": token_hash} async with client.tx() as tx: + await lock_provisioning_token(tx, token_hash) key: Final = await VerificationTokenRepository(SimpleNamespace(db=tx)).table.find_unique( where=LiteLLM_VerificationTokenWhereUniqueInput(token=token_hash) ) diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_source_endpoints.py b/tests/test_litellm/proxy/management_endpoints/scim/test_source_endpoints.py index 9de7f0c1110..c60cc2d9118 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_source_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_source_endpoints.py @@ -34,6 +34,7 @@ def source_database(monkeypatch: pytest.MonkeyPatch): client: Final = MagicMock(spec=PrismaClient) tx: Final = client.tx.return_value.__aenter__.return_value monkeypatch.setattr(proxy_server, "prisma_client", client) + tx.execute_raw = AsyncMock() tx.litellm_scimsource.find_unique = AsyncMock(return_value=None) tx.litellm_verificationtoken.find_unique = AsyncMock(return_value=SimpleNamespace(allowed_routes=["/scim/*"])) tx.litellm_accessgrouptable.find_many = AsyncMock(return_value=[]) @@ -239,3 +240,18 @@ async def test_source_mapping_accepts_access_groups_across_query_batches( ) assert result.group_mappings[0].access_group_ids == group_ids assert tx.litellm_accessgrouptable.find_many.await_count == 2 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("current_key", [None, SimpleNamespace(allowed_routes=["/*"])]) +async def test_source_creation_revalidates_key_after_waiting_for_mutation(monkeypatch, current_key): + tx = source_database(monkeypatch) + + async def complete_concurrent_mutation(*args): + tx.litellm_verificationtoken.find_unique.return_value = current_key + + tx.execute_raw.side_effect = complete_concurrent_mutation + with pytest.raises(HTTPException) as denied: + await create_source(SCIMSourceCreate(display_name="Source", tenant_id=TENANT, provisioning_token="test-token"), ADMIN) + assert denied.value.status_code == 400 + tx.litellm_scimsource.create.assert_not_awaited() 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 265d781a7e4..fad0f290462 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 @@ -12755,6 +12755,10 @@ def _make_regenerate_mock_prisma(): return_value=None ) mock_prisma_client.jsonify_object = MagicMock(side_effect=lambda data: data) + mock_prisma_client.tx = MagicMock() + mock_prisma_client.tx.return_value.__aenter__.return_value = mock_prisma_client.db + mock_prisma_client.db.litellm_scimsource.find_unique = AsyncMock(return_value=None) + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=_make_regenerate_existing_key()) return mock_prisma_client @@ -14823,6 +14827,11 @@ async def test_execute_virtual_key_regeneration_cache_invalidation_with_token_ha return_value=None ) mock_prisma_client.jsonify_object = MagicMock(side_effect=lambda data: data) + mock_prisma_client.tx = MagicMock() + mock_prisma_client.tx.return_value.__aenter__.return_value = mock_prisma_client.db + mock_prisma_client.db.litellm_scimsource.find_unique = AsyncMock(return_value=None) + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=existing_key) + mock_user_api_key_cache = MagicMock() mock_proxy_logging_obj = MagicMock() @@ -21154,3 +21163,79 @@ async def test_unbound_scim_key_can_still_regenerate(): ) assert result.key is not None client.db.litellm_verificationtoken.update.assert_awaited_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("routes", [None, [], ["/*"], ["/scim/*", "/key/info"]]) +async def test_key_update_endpoint_preserves_source_token_restriction(monkeypatch, routes): + from types import SimpleNamespace + + from litellm.proxy import proxy_server + from litellm.proxy.management_endpoints.key_management_endpoints import update_key_fn + + key = _make_regenerate_existing_key().model_copy(update={"allowed_routes": ["/scim/*"]}) + _wire_update_key_fn(monkeypatch, key) + client = proxy_server.prisma_client + client.writer_db.litellm_scimsource.find_unique = AsyncMock(return_value=SimpleNamespace(source_id="source")) + request = MagicMock() + request.query_params = {} + with pytest.raises(ProxyException) as denied: + await update_key_fn( + request=request, data=UpdateKeyRequest(key=key.token, allowed_routes=routes), + user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None, + ) + assert str(denied.value.code) == "409" + client.update_data.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_regeneration_rechecks_source_binding_before_transaction_writes(): + from types import SimpleNamespace + + from litellm.proxy.management_endpoints.key_management_endpoints import _execute_virtual_key_regeneration + + client = _make_regenerate_mock_prisma() + client.writer_db.litellm_scimsource.find_unique = AsyncMock(return_value=None) + tx = AsyncMock() + client.tx = MagicMock() + client.tx.return_value.__aenter__.return_value = tx + tx.litellm_scimsource.find_unique.return_value = SimpleNamespace(source_id="concurrent-source") + key = _make_regenerate_existing_key().model_copy(update={"allowed_routes": ["/scim/*"]}) + tx.litellm_verificationtoken.find_unique.return_value = key + with _patch_regenerate_side_effects(): + with pytest.raises(HTTPException) as denied: + await _execute_virtual_key_regeneration( + prisma_client=client, key_in_db=key, hashed_api_key=key.token, key="sk-original", data=None, + user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None, + user_api_key_cache=MagicMock(), proxy_logging_obj=MagicMock(), + ) + assert denied.value.status_code == 409 + tx.litellm_verificationtoken.update.assert_not_awaited() + tx.litellm_deletedverificationtoken.create_many.assert_not_awaited() + client.db.litellm_verificationtoken.update.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("bound,routes,missing,expected", [(True, ["/*"], False, 409), (True, ["/scim/*"], False, None), (False, [], False, None), (False, [], True, 409)]) +async def test_route_write_rechecks_binding_and_preserves_supported_updates(bound, routes, missing, expected): + from types import SimpleNamespace + + from litellm.proxy.management_endpoints.key_management_endpoints import _update_key_row_with_soft_budget + + client = _make_regenerate_mock_prisma() + key = _make_regenerate_existing_key().model_copy(update={"allowed_routes": ["/scim/*"]}) + tx = client.db + tx.litellm_scimsource.find_unique.return_value = SimpleNamespace(source_id="source") if bound else None + tx.litellm_verificationtoken.find_unique.return_value = None if missing else key + tx.litellm_verificationtoken.update.return_value = key.model_copy(update={"allowed_routes": routes}) + request = UpdateKeyRequest(key=key.token, allowed_routes=routes) + write = _update_key_row_with_soft_budget(client, key.token, request, {"allowed_routes": routes}, key, "admin") + if expected: + with pytest.raises(HTTPException) as denied: + await write + assert denied.value.status_code == expected + tx.litellm_verificationtoken.update.assert_not_awaited() + else: + response = await write + assert response["data"]["allowed_routes"] == routes + tx.litellm_budgettable.update.assert_not_awaited()