Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
87 changes: 69 additions & 18 deletions litellm/proxy/management_endpoints/key_management_endpoints.py
Original file line number Diff line number Diff line change
Expand Up @@ -463,6 +463,52 @@ def handle_key_type(data: GenerateKeyRequest, data_json: dict) -> dict:
_NON_ADMIN_SAFE_ALLOWED_ROUTES_PRESETS = frozenset({"llm_api_routes", "info_routes"})


def _validate_caller_can_change_key_ownership(
data: Optional[BaseModel],
existing_key_row: Any,
user_api_key_dict: UserAPIKeyAuth,
) -> None:
"""
Non-admin callers must not rebind a key's ``user_id`` to a different
user. The ``user_id`` on a verification token is what
``_return_user_api_key_auth_obj`` resolves against ``litellm_usertable``
to derive the request's role; a non-admin rebinding their own key's
``user_id`` to a ``PROXY_ADMIN`` row promotes themselves.

``/key/update`` already enforces this inline; ``/key/regenerate`` did
not. Sharing the check keeps both endpoints — and any future
regenerate-style endpoint — consistent.
"""
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value:
return
if data is None:
return
# Distinguish "user_id omitted" from "user_id explicitly set to None".
# Both leave ``getattr(data, 'user_id', None)`` at None, but only the
# explicit-null variant survives ``model_dump(exclude_unset=True)`` in
# ``prepare_key_update_data`` and writes NULL to the token row —
# detaching the key from its user and bypassing the user-row
# role check on subsequent requests.
fields_set = getattr(data, "model_fields_set", None) or set()
if "user_id" not in fields_set:
return
incoming_user_id = getattr(data, "user_id", None)
if incoming_user_id is None or incoming_user_id == "":
raise HTTPException(
status_code=403,
detail="Non-admin users cannot remove the user_id from a key.",
)
existing_user_id = getattr(existing_key_row, "user_id", None)
if incoming_user_id != existing_user_id:
raise HTTPException(
status_code=403,
detail=(
f"Non-admin caller is not allowed to rebind the key from "
f"user={existing_user_id} to user={incoming_user_id}"
),
)


def _check_allowed_routes_caller_permission(
allowed_routes: Optional[list],
user_api_key_dict: UserAPIKeyAuth,
Expand Down Expand Up @@ -2094,23 +2140,11 @@ async def _validate_update_key_data(
user_api_key_dict=user_api_key_dict,
)

# Prevent non-admin from removing user_id (setting to empty string) (LIT-1884)
if data.user_id is not None and data.user_id == "" and not _is_proxy_admin:
raise HTTPException(
status_code=403,
detail="Non-admin users cannot remove the user_id from a key.",
)

# sanity check - prevent non-proxy admin user from updating key to belong to a different user
if (
data.user_id is not None
and data.user_id != existing_key_row.user_id
and not _is_proxy_admin
):
raise HTTPException(
status_code=403,
detail=f"User={data.user_id} is not allowed to update key={data.key} to belong to user={existing_key_row.user_id}",
)
_validate_caller_can_change_key_ownership(
data=data,
existing_key_row=existing_key_row,
user_api_key_dict=user_api_key_dict,
)

common_key_access_checks(
user_api_key_dict=user_api_key_dict,
Expand Down Expand Up @@ -3935,6 +3969,13 @@ async def _execute_virtual_key_regeneration(
"""Generate new token, update DB, invalidate cache, and return response."""
from litellm.proxy.proxy_server import hash_token

# Mirror the /key/update ownership rebind guard. See helper docstring.
_validate_caller_can_change_key_ownership(
data=data,
existing_key_row=key_in_db,
user_api_key_dict=user_api_key_dict,
)

# Apply the same membership rule used on /key/update: when the caller
# asks to point the regenerated key at a different organization_id,
# require they are a member of (or proxy admin over) the target org.
Expand Down Expand Up @@ -4111,7 +4152,17 @@ async def regenerate_key_fn( # noqa: PLR0915
allow_safe_presets=True,
)

is_master_key_regeneration = data and data.new_master_key is not None
# Premium-gate bypass for master-key rotation must verify the
# caller actually holds the master key, not just that the request
# body has a ``new_master_key`` field. A presence-only check let
# any non-premium caller skip the enterprise gate by sending any
# value in that field.
regenerate_target_key = data.key if data and data.key else key
is_master_key_regeneration = (
data is not None
and data.new_master_key is not None
and _is_master_key(api_key=regenerate_target_key, _master_key=master_key)
)

if (
premium_user is not True and not is_master_key_regeneration
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -10019,3 +10019,157 @@ def __iter__(self):
call_kwargs = mock_delete_cache.call_args.kwargs
# The token hash should be passed as-is, NOT double-hashed
assert call_kwargs["hashed_token"] == token_hash


# ---------------------------------------------------------------------------
# /key/regenerate ownership-rebind guard + premium-gate identity check
# ---------------------------------------------------------------------------

import contextlib # noqa: E402


def _non_admin_user_api_key_dict():
return UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER,
api_key="sk-internal",
user_id="user-1",
)


@contextlib.contextmanager
def _patch_regenerate_side_effects():
"""Mock out token creation + DB write + cache + rotation hook so
``_execute_virtual_key_regeneration`` runs to completion under test."""
with (
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.get_new_token",
new_callable=AsyncMock,
return_value="sk-newtoken1234ab12",
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_rotated_hook",
new_callable=AsyncMock,
),
):
yield


@pytest.mark.asyncio
@pytest.mark.parametrize(
"incoming_user_id,expected_status,expected_substring",
[
# Cross-user rebind: the privesc primitive.
("default_user_id", 403, "not allowed to rebind the key"),
# Empty-string removal: companion guard.
("", 403, "remove the user_id"),
# Explicit null: same effect as empty-string removal — survives
# model_dump(exclude_unset=True) and writes NULL to the token row.
(None, 403, "remove the user_id"),
# No-op rebind (caller sends their own user_id): must succeed.
("user-1", None, None),
],
ids=[
"rebind_blocked",
"empty_blocked",
"explicit_null_blocked",
"same_user_id_allowed",
],
)
async def test_regenerate_user_id_rebind_guard(
incoming_user_id, expected_status, expected_substring
):
from litellm.proxy._types import RegenerateKeyRequest
from litellm.proxy.management_endpoints.key_management_endpoints import (
_execute_virtual_key_regeneration,
)

existing_key = _make_regenerate_existing_key()
data = RegenerateKeyRequest(user_id=incoming_user_id)

async def _run():
await _execute_virtual_key_regeneration(
prisma_client=_make_regenerate_mock_prisma(),
key_in_db=existing_key,
hashed_api_key="abc123",
key="abc123",
data=data,
user_api_key_dict=_non_admin_user_api_key_dict(),
litellm_changed_by=None,
user_api_key_cache=MagicMock(),
proxy_logging_obj=MagicMock(),
)

if expected_status is None:
with _patch_regenerate_side_effects():
await _run()
return

with pytest.raises(HTTPException) as exc:
await _run()
assert exc.value.status_code == expected_status
assert expected_substring in str(exc.value.detail)


@pytest.mark.asyncio
async def test_regenerate_premium_gate_requires_actual_master_key():
# ``regenerate_key_fn``'s decorator wraps the underlying ValueError
# into a ProxyException with empty ``message``. The exception type
# alone confirms the premium gate fired.
from litellm.proxy._types import RegenerateKeyRequest
from litellm.proxy.management_endpoints.key_management_endpoints import (
regenerate_key_fn,
)

data = RegenerateKeyRequest(key="sk-not-master", new_master_key="anything")

with (
patch("litellm.proxy.proxy_server.premium_user", False),
patch("litellm.proxy.proxy_server.master_key", "sk-the-real-master-key"),
patch("litellm.proxy.proxy_server.prisma_client", AsyncMock()),
pytest.raises((ValueError, HTTPException, ProxyException)),
):
await regenerate_key_fn(
key="sk-not-master",
data=data,
user_api_key_dict=_non_admin_user_api_key_dict(),
)


@pytest.mark.asyncio
async def test_regenerate_premium_gate_allows_actual_master_key_holder():
from litellm.proxy._types import RegenerateKeyRequest
from litellm.proxy.management_endpoints.key_management_endpoints import (
regenerate_key_fn,
)

master = "sk-the-real-master-key"
data = RegenerateKeyRequest(key=master, new_master_key="sk-new-master")

with (
patch("litellm.proxy.proxy_server.premium_user", False),
patch("litellm.proxy.proxy_server.master_key", master),
patch("litellm.proxy.proxy_server.prisma_client", AsyncMock()),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._rotate_master_key",
new_callable=AsyncMock,
),
):
result = await regenerate_key_fn(
key=master,
data=data,
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key=master,
user_id="admin",
),
)

assert result.token == "sk-new-master"
Loading