mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(scim): serialize source binding with token mutations
This commit is contained in:
parent
d8efd6a151
commit
d362a91d2b
4 changed files with 168 additions and 31 deletions
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue