fix(scim): serialize source binding with token mutations

This commit is contained in:
Joshua Valluru 2026-09-28 21:52:40 -07:00
parent d8efd6a151
commit d362a91d2b
4 changed files with 168 additions and 31 deletions

View file

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

View file

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

View file

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

View file

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