fix(scim): guard source token regeneration and route changes

This commit is contained in:
Joshua Valluru 2026-09-28 19:37:23 -07:00
parent ee9d32ccaf
commit d8efd6a151
2 changed files with 79 additions and 0 deletions

View file

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

View file

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