fix: use standard patterns, make test more robust

This commit is contained in:
Ryan Crabbe 2026-03-09 15:29:24 -07:00
parent d761501ce2
commit a23ab1972c
2 changed files with 42 additions and 17 deletions

View file

@ -1168,11 +1168,10 @@ async def generate_key_fn(
# Auto-populate user_id from authenticated user when not provided (non-admins only)
if data.user_id is None and user_api_key_dict.user_id is not None:
if (
user_api_key_dict.user_role is None
or user_api_key_dict.user_role
!= LitellmUserRoles.PROXY_ADMIN.value
):
if user_api_key_dict.user_role not in [
LitellmUserRoles.PROXY_ADMIN.value,
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value,
]:
data.user_id = user_api_key_dict.user_id
# Validate budget values are not negative
@ -1894,7 +1893,11 @@ async def update_key_fn(
"user_id" in _update_fields
and _update_fields["user_id"] is None
and existing_key_row.user_id is not None
and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value
and user_api_key_dict.user_role
not in [
LitellmUserRoles.PROXY_ADMIN.value,
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value,
]
):
raise HTTPException(
status_code=403,

View file

@ -6831,15 +6831,37 @@ async def test_generate_key_validation_user_id_and_team_id(monkeypatch):
)
assert exc.value.code == "404"
# 3. Non-admin cannot remove user_id on update (guard logic)
update_request = UpdateKeyRequest(key="sk-1", user_id=None)
update_request.model_fields_set.add("user_id")
_update_fields = update_request.model_dump(exclude_unset=True)
existing_key = LiteLLM_VerificationToken(token="t", user_id="original-user")
is_blocked = (
"user_id" in _update_fields
and _update_fields["user_id"] is None
and existing_key.user_id is not None
and LitellmUserRoles.INTERNAL_USER.value != LitellmUserRoles.PROXY_ADMIN.value
# 3. Non-admin cannot remove user_id on update (calls update_key_fn)
from starlette.requests import Request as StarletteRequest
from starlette.datastructures import Headers
from litellm.proxy.management_endpoints.key_management_endpoints import (
update_key_fn,
)
assert is_blocked is True
# Mock prisma_client.get_data to return an existing key with user_id set
mock_prisma_client.get_data = AsyncMock(
return_value=LiteLLM_VerificationToken(
token="hashed_sk1", user_id="original-user"
)
)
# Build a minimal ASGI request
scope = {"type": "http", "method": "POST", "headers": [], "query_string": b""}
mock_request = StarletteRequest(scope)
# UpdateKeyRequest with explicit user_id=None
update_data = UpdateKeyRequest(key="sk-1", user_id=None)
update_data.model_fields_set.add("user_id")
with pytest.raises(ProxyException) as exc:
await update_key_fn(
request=mock_request,
data=update_data,
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER,
api_key="sk-user",
user_id="original-user",
),
)
assert exc.value.code == "403"