mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-21 00:21:49 +00:00
fix(proxy): validate bulk object_permission against the key's team as /key/update does
This commit is contained in:
parent
f982d3e046
commit
df6a222cb8
2 changed files with 78 additions and 14 deletions
|
|
@ -2798,9 +2798,18 @@ async def _process_single_key_update(
|
|||
llm_router=llm_router,
|
||||
)
|
||||
|
||||
key_request: Final = await _with_validated_object_permission(
|
||||
update_key_request=update_key_request,
|
||||
team_obj=team_obj,
|
||||
existing_key_row=existing_key_row,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
# Prepare update data
|
||||
non_default_values = await prepare_key_update_data(
|
||||
data=update_key_request, existing_key_row=existing_key_row, prisma_client=prisma_client, llm_router=llm_router
|
||||
data=key_request, existing_key_row=existing_key_row, prisma_client=prisma_client, llm_router=llm_router
|
||||
)
|
||||
|
||||
await _enforce_custom_key_policy(
|
||||
|
|
@ -2809,7 +2818,7 @@ async def _process_single_key_update(
|
|||
operation="update",
|
||||
existing_key_row=existing_key_row,
|
||||
non_default_values=non_default_values,
|
||||
request=update_key_request,
|
||||
request=key_request,
|
||||
),
|
||||
)
|
||||
|
||||
|
|
@ -2825,15 +2834,15 @@ async def _process_single_key_update(
|
|||
existing_key_row=existing_key_row,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
_data: Final = {**update_values, "token": update_key_request.key}
|
||||
_data: Final = {**update_values, "token": key_request.key}
|
||||
response: Final[Mapping[str, object] | None] = cast( # cast-ok: every update_data branch returns a str-keyed dict
|
||||
"Mapping[str, object] | None",
|
||||
await prisma_client.update_data(token=update_key_request.key, data=_data),
|
||||
await prisma_client.update_data(token=key_request.key, data=_data),
|
||||
)
|
||||
|
||||
# Delete cache
|
||||
await _delete_cache_key_object(
|
||||
hashed_token=_hash_token_if_needed(update_key_request.key),
|
||||
hashed_token=_hash_token_if_needed(key_request.key),
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
|
@ -2842,17 +2851,15 @@ async def _process_single_key_update(
|
|||
# authenticating against the access groups it just lost.
|
||||
await sync_key_update_access_group_membership(
|
||||
prisma_client=prisma_client,
|
||||
key_token=_hash_token_if_needed(
|
||||
_resolve_token_to_update(data=update_key_request, existing_key_row=existing_key_row)
|
||||
),
|
||||
data=update_key_request,
|
||||
key_token=_hash_token_if_needed(_resolve_token_to_update(data=key_request, existing_key_row=existing_key_row)),
|
||||
data=key_request,
|
||||
existing_key_row=existing_key_row,
|
||||
)
|
||||
|
||||
# Trigger async hook
|
||||
asyncio.create_task(
|
||||
KeyManagementEventHooks.async_key_updated_hook(
|
||||
data=update_key_request,
|
||||
data=key_request,
|
||||
existing_key_row=existing_key_row,
|
||||
response=response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
|
|
@ -2875,6 +2882,31 @@ async def _process_single_key_update(
|
|||
return updated_key_info
|
||||
|
||||
|
||||
async def _with_validated_object_permission(
|
||||
update_key_request: UpdateKeyRequest,
|
||||
team_obj: LiteLLM_TeamTableCachedObj | None,
|
||||
existing_key_row: LiteLLM_VerificationToken,
|
||||
prisma_client: PrismaClient | None,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> UpdateKeyRequest:
|
||||
if update_key_request.object_permission is None:
|
||||
return update_key_request
|
||||
normalized_object_permission: Final = await _validate_mcp_servers_for_key_update(
|
||||
data=update_key_request,
|
||||
team_obj=team_obj,
|
||||
existing_key_row=existing_key_row,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
is_proxy_admin=user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value,
|
||||
)
|
||||
if normalized_object_permission is None:
|
||||
return update_key_request
|
||||
return update_key_request.model_copy(
|
||||
update=MappingProxyType({"object_permission": LiteLLM_ObjectPermissionBase(**normalized_object_permission)})
|
||||
)
|
||||
|
||||
|
||||
async def _validate_mcp_servers_for_key_update(
|
||||
data: "UpdateKeyRequest",
|
||||
team_obj: Optional["LiteLLM_TeamTableCachedObj"],
|
||||
|
|
|
|||
|
|
@ -80,7 +80,11 @@ from litellm.proxy.management_endpoints.key_management_endpoints import (
|
|||
validate_key_team_change,
|
||||
)
|
||||
from litellm.proxy.proxy_server import app
|
||||
from litellm.types.proxy.management_endpoints.key_management_endpoints import CustomKeyPolicyRequest
|
||||
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
|
||||
BulkUpdateKeyRequest,
|
||||
BulkUpdateKeyResponse,
|
||||
CustomKeyPolicyRequest,
|
||||
)
|
||||
|
||||
client = TestClient(app)
|
||||
|
||||
|
|
@ -7098,23 +7102,29 @@ async def test_list_key_helper_applies_search_to_prisma_where():
|
|||
|
||||
|
||||
_BULK_UPDATE_TOKEN: Final = "1f2e3d4c5b6a79880123456789abcdef0123456789abcdef0123456789abcdef"
|
||||
_BULK_UPDATE_TEAM: Final = LiteLLM_TeamTableCachedObj(team_id="team-1")
|
||||
|
||||
|
||||
async def _bulk_update_one_key(monkeypatch, item_payload: Mapping[str, object]) -> AsyncMock:
|
||||
async def _run_bulk_update_on_one_key(
|
||||
monkeypatch, item_payload: Mapping[str, object], team: LiteLLM_TeamTableCachedObj = _BULK_UPDATE_TEAM
|
||||
) -> tuple[BulkUpdateKeyResponse, AsyncMock]:
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import bulk_update_keys
|
||||
from litellm.types.proxy.management_endpoints.key_management_endpoints import BulkUpdateKeyRequest
|
||||
|
||||
key_in_db = LiteLLM_VerificationToken(
|
||||
token=_BULK_UPDATE_TOKEN, user_id="test-user", team_id="team-1", max_budget=100.0, budget_id="budget-1"
|
||||
)
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=key_in_db)
|
||||
mock_prisma_client.get_data = AsyncMock(return_value=key_in_db)
|
||||
mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=None)
|
||||
mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock(
|
||||
return_value=MagicMock(object_permission_id="objperm-bulk")
|
||||
)
|
||||
mock_prisma_client.update_data = AsyncMock(return_value={"data": {"token": _BULK_UPDATE_TOKEN}})
|
||||
_setup_update_key_mocks(monkeypatch, mock_prisma_client)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", AsyncMock(return_value=team)
|
||||
)
|
||||
|
||||
with (
|
||||
patch( # test-quality-ok: the handler reads the cache and hook singletons from module globals, no injection seam
|
||||
|
|
@ -7138,8 +7148,13 @@ async def _bulk_update_one_key(monkeypatch, item_payload: Mapping[str, object])
|
|||
litellm_changed_by=None,
|
||||
)
|
||||
|
||||
return response, mock_prisma_client
|
||||
|
||||
|
||||
async def _bulk_update_one_key(monkeypatch, item_payload: Mapping[str, object]) -> AsyncMock:
|
||||
response, prisma = await _run_bulk_update_on_one_key(monkeypatch, item_payload)
|
||||
assert response.failed_updates == []
|
||||
return mock_prisma_client
|
||||
return prisma
|
||||
|
||||
|
||||
def _written_key_row(prisma: AsyncMock) -> Mapping[str, object]:
|
||||
|
|
@ -7178,6 +7193,23 @@ async def test_bulk_update_keys_object_permission_is_granted_not_dropped(monkeyp
|
|||
assert not {"max_budget", "team_id", "budget_id"} & written.keys()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bulk_update_keys_object_permission_outside_the_team_allowlist_is_refused(monkeypatch):
|
||||
"""A bulk item's object_permission is checked against the key's team exactly as /key/update
|
||||
checks it, so a team key cannot be granted a search tool its team does not allow."""
|
||||
team = LiteLLM_TeamTableCachedObj(
|
||||
team_id="team-1",
|
||||
object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="op-team-1", search_tools=["team-search"]),
|
||||
)
|
||||
response, prisma = await _run_bulk_update_on_one_key(
|
||||
monkeypatch, {"object_permission": {"search_tools": ["other-search"]}}, team=team
|
||||
)
|
||||
|
||||
assert response.successful_updates == []
|
||||
assert "not allowed by team 'team-1'" in response.failed_updates[0].failed_reason
|
||||
prisma.update_data.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generate_key_negative_max_budget():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue