From d8efd6a151c8c5d36f2d84da469e9bf2c60f31a5 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 19:37:23 -0700 Subject: [PATCH] fix(scim): guard source token regeneration and route changes --- .../key_management_endpoints.py | 22 +++++++ .../test_key_management_endpoints.py | 57 +++++++++++++++++++ 2 files changed, 79 insertions(+) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 2de9ddc2577..a8124ac6a12 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -142,6 +142,7 @@ from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import ( DeletedVerificationTokenRepository, DeprecatedVerificationTokenRepository, + SCIMSourceRepository, ) from litellm.repositories.team_repository import TeamRepository from litellm.repositories.user_repository import UserRepository @@ -2744,6 +2745,13 @@ async def _process_single_key_update( prisma_client=prisma_client, ) + if ( + "allowed_routes" in update_key_request.model_fields_set + and tuple(update_key_request.allowed_routes or ()) != ("/scim/*",) + and prisma_client is not None + ): + await _reject_source_bound_key_change(prisma_client, existing_key_row) + _check_disable_global_guardrails_caller_permission( update_key_request.disable_global_guardrails, update_key_request.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # request models declare `metadata` as bare dict @@ -5532,6 +5540,18 @@ async def _insert_deprecated_key( ) +async def _reject_source_bound_key_change(prisma_client: PrismaClient, key: LiteLLM_VerificationToken) -> None: + if tuple(key.allowed_routes or ()) != ("/scim/*",): + return + source: Final = await SCIMSourceRepository(prisma_client, use_writer=True).table.find_unique( + where={"key_hash": key.token} + ) + if source is not None: + raise HTTPException( + 409, "A provisioning source token cannot be regenerated or have its SCIM restriction removed" + ) + + async def _execute_virtual_key_regeneration( *, prisma_client: PrismaClient, @@ -5549,6 +5569,8 @@ async def _execute_virtual_key_regeneration( from litellm.proxy import proxy_server from litellm.proxy.proxy_server import hash_token + await _reject_source_bound_key_change(prisma_client, key_in_db) + # Mirror the /key/update ownership rebind guard. See helper docstring. _validate_caller_can_change_key_ownership( data=data, 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 5ea38ce23d5..265d781a7e4 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 @@ -21097,3 +21097,60 @@ class TestTeamAdminMemberKeyBudgetUpdate: ) assert exc.value.status_code == 403 assert "member_key_budgets" not in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_source_bound_key_cannot_regenerate_before_any_write(): + 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=SimpleNamespace(source_id="source")) + key = _make_regenerate_existing_key().model_copy(update={"allowed_routes": ["/scim/*"]}) + 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 + client.db.litellm_verificationtoken.update.assert_not_awaited() + client.db.litellm_verificationtoken.create.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("routes", [None, [], ["llm_api_routes"], ["/scim/*", "/key/info"]]) +async def test_source_bound_key_cannot_remove_its_route_restriction(routes): + from types import SimpleNamespace + from litellm.proxy.management_endpoints.key_management_endpoints import _process_single_key_update + + client = _make_regenerate_mock_prisma() + client.update_data = AsyncMock(return_value={"data": {"key_alias": "source"}}) + client.writer_db.litellm_scimsource.find_unique = AsyncMock(return_value=SimpleNamespace(source_id="source")) + key = _make_regenerate_existing_key().model_copy(update={"allowed_routes": ["/scim/*"]}) + with pytest.raises(HTTPException) as denied: + await _process_single_key_update( + update_key_request=UpdateKeyRequest(key=key.token, allowed_routes=routes), existing_key_row=key, + user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None, + prisma_client=client, user_api_key_cache=MagicMock(), proxy_logging_obj=MagicMock(), llm_router=None, + ) + assert denied.value.status_code == 409 + client.update_data.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_unbound_scim_key_can_still_regenerate(): + 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) + key = _make_regenerate_existing_key().model_copy(update={"allowed_routes": ["/scim/*"]}) + with _patch_regenerate_side_effects(): + result = 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 result.key is not None + client.db.litellm_verificationtoken.update.assert_awaited_once()