diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py index 92c8dc344ae..173acd41b35 100644 --- a/litellm/proxy/management_endpoints/scim/scim_v2.py +++ b/litellm/proxy/management_endpoints/scim/scim_v2.py @@ -337,41 +337,72 @@ async def _handle_team_membership_changes( ) +SCIM_BLOCKED_METADATA_KEY = "scim_blocked" + + +def _key_was_scim_blocked(metadata: Any) -> bool: + """True if a verification token carries the SCIM-block marker in metadata.""" + return ( + isinstance(metadata, dict) and metadata.get(SCIM_BLOCKED_METADATA_KEY) is True + ) + + async def _set_user_keys_blocked(user_id: str, blocked: bool) -> int: """ - Block or unblock all virtual keys owned by a user and invalidate them in - the in-memory/redis caches so the change takes effect immediately. + Block or unblock virtual keys owned by a user and invalidate them in the + in-memory/redis caches so the change takes effect immediately. - Returns the number of keys whose state was flipped. Used by the SCIM - deprovisioning flow so a user's keys stop working the moment SCIM marks - the user inactive (or deletes them). + Each key SCIM blocks is tagged with ``metadata.scim_blocked = True``. On + reactivation we only unblock keys carrying that marker, so a key an admin + blocked manually for unrelated reasons is left alone. + + Returns the number of keys whose state was flipped. """ from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache prisma_client = await _get_prisma_client_or_raise_exception() - # Only flip keys whose current state differs — avoids touching keys that - # were already (un)blocked manually by an admin. `blocked` is a nullable - # column with no default, so existing keys typically have `blocked=NULL`; - # we must treat NULL as "not blocked" so SQL equality on NULL doesn't - # silently skip them. if blocked: - state_filter: Dict[str, Any] = {"OR": [{"blocked": False}, {"blocked": None}]} + # Block keys that aren't already blocked. `blocked` is a nullable column + # with no default so existing rows typically hold NULL; treat NULL as + # "not blocked" so SQL equality on NULL doesn't silently skip them. + candidates = await prisma_client.db.litellm_verificationtoken.find_many( + where={ + "user_id": user_id, + "OR": [{"blocked": False}, {"blocked": None}], + }, + ) + affected_keys = candidates else: - state_filter = {"blocked": True} + # Only unblock keys that SCIM previously blocked. An admin-managed + # block has no `scim_blocked` marker and must not be reversed here. + candidates = await prisma_client.db.litellm_verificationtoken.find_many( + where={"user_id": user_id, "blocked": True}, + ) + affected_keys = [k for k in candidates if _key_was_scim_blocked(k.metadata)] - where_clause: Dict[str, Any] = {"user_id": user_id, **state_filter} - - affected_keys = await prisma_client.db.litellm_verificationtoken.find_many( - where=where_clause, - ) if not affected_keys: return 0 - await prisma_client.db.litellm_verificationtoken.update_many( - where=where_clause, - data={"blocked": blocked}, - ) + # Per-key updates: we need to add/remove the SCIM-block marker in JSON + # metadata, which `update_many` can't express. Cardinality is bounded by + # the number of keys a single user owns. + for key_row in affected_keys: + current_metadata: Dict[str, Any] = ( + dict(key_row.metadata) if isinstance(key_row.metadata, dict) else {} + ) + if blocked: + new_metadata = {**current_metadata, SCIM_BLOCKED_METADATA_KEY: True} + else: + new_metadata = { + k: v + for k, v in current_metadata.items() + if k != SCIM_BLOCKED_METADATA_KEY + } + await prisma_client.db.litellm_verificationtoken.update( + where={"token": key_row.token}, + data={"blocked": blocked, "metadata": safe_dumps(new_metadata)}, + ) for key_row in affected_keys: await _delete_cache_key_object( diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_key_deactivation.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_key_deactivation.py index e68aef9aeda..e14209ce068 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_key_deactivation.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_key_deactivation.py @@ -24,11 +24,12 @@ from litellm.types.proxy.management_endpoints.scim_v2 import ( ) -def _build_token_row(token: str, user_id: str, blocked: bool): +def _build_token_row(token: str, user_id: str, blocked: bool, metadata=None): row = MagicMock() row.token = token row.user_id = user_id row.blocked = blocked + row.metadata = metadata if metadata is not None else {} return row @@ -45,6 +46,7 @@ def _build_prisma_with_keys(user_keys, mock_user=None, updated_user=None): mock_db.litellm_verificationtoken.find_many = AsyncMock(return_value=user_keys) mock_db.litellm_verificationtoken.update_many = AsyncMock(return_value=None) + mock_db.litellm_verificationtoken.update = AsyncMock(return_value=None) return mock_client, mock_db @@ -74,13 +76,17 @@ async def test_set_user_keys_blocked_flips_state_and_invalidates_cache(): flipped = await _set_user_keys_blocked(user_id="user-x", blocked=True) assert flipped == 2 - mock_db.litellm_verificationtoken.update_many.assert_awaited_once_with( - where={ - "user_id": "user-x", - "OR": [{"blocked": False}, {"blocked": None}], - }, - data={"blocked": True}, - ) + # Each key gets a per-row update that flips `blocked` and stamps the + # SCIM-block marker into metadata. + assert mock_db.litellm_verificationtoken.update.await_count == 2 + update_calls = mock_db.litellm_verificationtoken.update.await_args_list + seen_tokens = set() + for call in update_calls: + kwargs = call.kwargs or call[1] + seen_tokens.add(kwargs["where"]["token"]) + assert kwargs["data"]["blocked"] is True + assert '"scim_blocked": true' in kwargs["data"]["metadata"] + assert seen_tokens == {"hash-1", "hash-2"} assert sorted(cache_deletions) == ["hash-1", "hash-2"] @@ -101,10 +107,48 @@ async def test_set_user_keys_blocked_noop_when_no_matching_keys(): flipped = await _set_user_keys_blocked(user_id="user-x", blocked=True) assert flipped == 0 + mock_db.litellm_verificationtoken.update.assert_not_called() mock_db.litellm_verificationtoken.update_many.assert_not_called() mocked_delete.assert_not_called() +@pytest.mark.asyncio +async def test_set_user_keys_unblocked_skips_admin_blocked_keys(): + """Reactivation must leave keys an admin blocked (no scim_blocked marker) alone.""" + keys = [ + # SCIM-blocked: should be unblocked. + _build_token_row( + "hash-scim", "user-x", blocked=True, metadata={"scim_blocked": True} + ), + # Admin-blocked for unrelated reasons: must remain blocked. + _build_token_row("hash-admin", "user-x", blocked=True, metadata={}), + ] + mock_client, mock_db = _build_prisma_with_keys(keys) + + cache_deletions = [] + + async def fake_delete(hashed_token, user_api_key_cache, proxy_logging_obj): + cache_deletions.append(hashed_token) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_client), + patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + patch( + "litellm.proxy.management_endpoints.scim.scim_v2._delete_cache_key_object", + AsyncMock(side_effect=fake_delete), + ), + ): + flipped = await _set_user_keys_blocked(user_id="user-x", blocked=False) + + assert flipped == 1 + mock_db.litellm_verificationtoken.update.assert_awaited_once() + update_kwargs = mock_db.litellm_verificationtoken.update.await_args.kwargs + assert update_kwargs["where"] == {"token": "hash-scim"} + assert update_kwargs["data"]["blocked"] is False + assert cache_deletions == ["hash-scim"] + + @pytest.mark.asyncio async def test_scim_delete_user_blocks_keys_before_deleting_user(): """SCIM DELETE /Users/{id} must block the user's keys before removing the row.""" @@ -131,13 +175,11 @@ async def test_scim_delete_user_blocks_keys_before_deleting_user(): response = await delete_user(user_id=user_id) assert response.status_code == 204 - mock_db.litellm_verificationtoken.update_many.assert_awaited_once_with( - where={ - "user_id": user_id, - "OR": [{"blocked": False}, {"blocked": None}], - }, - data={"blocked": True}, - ) + mock_db.litellm_verificationtoken.update.assert_awaited_once() + update_kwargs = mock_db.litellm_verificationtoken.update.await_args.kwargs + assert update_kwargs["where"] == {"token": "hash-a"} + assert update_kwargs["data"]["blocked"] is True + assert '"scim_blocked": true' in update_kwargs["data"]["metadata"] mock_db.litellm_usertable.delete.assert_awaited_once_with( where={"user_id": user_id} ) @@ -192,13 +234,11 @@ async def test_scim_patch_user_active_false_blocks_keys(): ): await patch_user(user_id=user_id, patch_ops=patch_ops) - mock_db.litellm_verificationtoken.update_many.assert_awaited_once_with( - where={ - "user_id": user_id, - "OR": [{"blocked": False}, {"blocked": None}], - }, - data={"blocked": True}, - ) + mock_db.litellm_verificationtoken.update.assert_awaited_once() + update_kwargs = mock_db.litellm_verificationtoken.update.await_args.kwargs + assert update_kwargs["where"] == {"token": "hash-z"} + assert update_kwargs["data"]["blocked"] is True + assert '"scim_blocked": true' in update_kwargs["data"]["metadata"] @pytest.mark.asyncio @@ -218,7 +258,11 @@ async def test_scim_patch_user_active_true_unblocks_keys(): teams=[], metadata={"scim_active": True, "scim_metadata": {}}, ) - keys = [_build_token_row("hash-r", user_id, blocked=True)] + keys = [ + _build_token_row( + "hash-r", user_id, blocked=True, metadata={"scim_blocked": True} + ) + ] mock_client, mock_db = _build_prisma_with_keys( keys, mock_user=mock_user, updated_user=updated_user ) @@ -250,10 +294,12 @@ async def test_scim_patch_user_active_true_unblocks_keys(): ): await patch_user(user_id=user_id, patch_ops=patch_ops) - mock_db.litellm_verificationtoken.update_many.assert_awaited_once_with( - where={"user_id": user_id, "blocked": True}, - data={"blocked": False}, - ) + mock_db.litellm_verificationtoken.update.assert_awaited_once() + update_kwargs = mock_db.litellm_verificationtoken.update.await_args.kwargs + assert update_kwargs["where"] == {"token": "hash-r"} + assert update_kwargs["data"]["blocked"] is False + # SCIM-block marker is stripped on reactivation. + assert "scim_blocked" not in update_kwargs["data"]["metadata"] @pytest.mark.asyncio