mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(scim): guard source token regeneration and route changes
This commit is contained in:
parent
ee9d32ccaf
commit
d8efd6a151
2 changed files with 79 additions and 0 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue