fix(scim): preserve admin-blocked keys across SCIM reactivation

Tag each key SCIM blocks with metadata.scim_blocked=True. On reactivation
unblock only those keys, leaving keys an admin blocked for unrelated
reasons untouched.
This commit is contained in:
mateo-berri 2026-05-01 19:34:56 +00:00 • committed by Cursor Agent
parent dc123d9f12
commit 8c409006ad
No known key found for this signature in database
2 changed files with 125 additions and 48 deletions

View file

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

View file

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