test(proxy): type user endpoint test doubles

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
mrinal 2026-10-02 01:08:46 +00:00
parent 88195a49ce
commit f552e86fea

View file

@ -5,7 +5,7 @@ import logging
from collections.abc import Mapping, Sequence
from datetime import datetime, timezone
from types import SimpleNamespace
from typing import Final
from typing import TYPE_CHECKING, Final, cast
from unittest.mock import AsyncMock, MagicMock
import httpx
@ -45,6 +45,10 @@ from tests.unit.proxy.management_endpoints.jwt_key_mapping_doubles import (
client = TestClient(app)
if TYPE_CHECKING:
from litellm.caching.redis_cache import RedisCache
from litellm.proxy.utils import PrismaClient
@pytest.mark.asyncio
async def test_ui_view_users_with_null_email(mocker, caplog):
@ -2362,43 +2366,79 @@ async def test_user_max_budget_update_evicts_cached_user_on_every_worker(mocker:
async def test_user_rate_limit_update_reaches_cached_user_on_every_worker(
mocker: MockerFixture, field: str, new_limit: int | None, all_users: bool
) -> None:
from redis.asyncio import Redis
from litellm.proxy._types import LiteLLM_UserTable
from litellm.proxy.auth.auth_checks import get_user_object
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import AuthCacheInvalidationSubscriber
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.management_endpoints.internal_user_endpoints import bulk_user_update, user_update
from litellm.types.proxy.management_endpoints.internal_user_endpoints import BulkUpdateUserRequest
from litellm.types.proxy.management_endpoints.internal_user_endpoints import (
BulkUpdateUserRequest,
UpdateUserRequestNoUserIDorEmail,
)
published: Final[list[tuple[str, str]]] = [] # mutable-ok: captures messages from the async Redis publisher
class _RecordingRedisClient(Redis):
def __init__(self) -> None:
pass
class _RecordingRedisClient:
async def publish(self, channel: str, message: str) -> int:
published.append((channel, message))
return 1
class _FakeRedisCache:
def __init__(self) -> None:
self.namespace = None
namespace: str | None = None
def init_pubsub_client(self) -> object:
def init_pubsub_client(self) -> _RecordingRedisClient:
return _RecordingRedisClient()
class _UserTableMocks:
def __init__(
self,
find_first: AsyncMock,
find_unique: AsyncMock,
find_many: AsyncMock,
update_many: AsyncMock,
) -> None:
self.find_first = find_first
self.find_unique = find_unique
self.find_many = find_many
self.update_many = update_many
class _DatabaseMocks:
def __init__(self, litellm_usertable: _UserTableMocks) -> None:
self.litellm_usertable = litellm_usertable
class _PrismaClientMock:
def __init__(
self,
db: _DatabaseMocks,
get_data: AsyncMock,
updated_user: LiteLLM_UserTable,
) -> None:
self.db = db
self.get_data = get_data
self.updated_user = updated_user
self.update_data_payload: dict[str, object] | None = None
async def update_data(self, user_id: str, data: dict[str, object], table_name: str) -> dict[str, object]:
self.update_data_payload = data
return {"user_id": user_id, "data": self.updated_user}
saved_user: Final = LiteLLM_UserTable(user_id="user-spruce", tpm_limit=100000, rpm_limit=1000)
updated_user: Final = saved_user.model_copy(update={field: new_limit})
old_limit: Final = 100000 if field == "tpm_limit" else 1000
prisma_client: Final = mocker.MagicMock()
prisma_client.db.litellm_usertable.find_first = mocker.AsyncMock(return_value=saved_user)
prisma_client.db.litellm_usertable.find_unique = mocker.AsyncMock(return_value=updated_user)
prisma_client.db.litellm_usertable.find_many = mocker.AsyncMock(return_value=[saved_user])
prisma_client.db.litellm_usertable.update_many = mocker.AsyncMock(return_value=1)
prisma_client.get_data = mocker.AsyncMock(return_value=saved_user)
prisma_client.update_data = mocker.AsyncMock(return_value={"user_id": saved_user.user_id, "data": updated_user})
prisma_client: Final = _PrismaClientMock(
db=_DatabaseMocks(
litellm_usertable=_UserTableMocks(
find_first=mocker.AsyncMock(return_value=saved_user),
find_unique=mocker.AsyncMock(return_value=updated_user),
find_many=mocker.AsyncMock(return_value=[saved_user]),
update_many=mocker.AsyncMock(return_value=1),
)
),
get_data=mocker.AsyncMock(return_value=saved_user),
updated_user=updated_user,
)
prisma_client_for_auth: Final = cast("PrismaClient", prisma_client)
mocker.patch( # test-quality-ok: substitute the database dependency
"litellm.proxy.proxy_server.prisma_client", prisma_client
)
@ -2425,13 +2465,13 @@ async def test_user_rate_limit_update_reaches_cached_user_on_every_worker(
handling_user_before: Final = await get_user_object(
user_id=saved_user.user_id,
prisma_client=prisma_client,
prisma_client=prisma_client_for_auth,
user_api_key_cache=handling_worker_cache,
user_id_upsert=False,
)
other_user_before: Final = await get_user_object(
user_id=saved_user.user_id,
prisma_client=prisma_client,
prisma_client=prisma_client_for_auth,
user_api_key_cache=other_worker_cache,
user_id_upsert=False,
)
@ -2443,22 +2483,26 @@ async def test_user_rate_limit_update_reaches_cached_user_on_every_worker(
admin: Final = UserAPIKeyAuth(user_id="admin-spruce", user_role=LitellmUserRoles.PROXY_ADMIN)
if all_users:
await bulk_user_update(
data=BulkUpdateUserRequest(all_users=True, user_updates={field: new_limit}),
data=BulkUpdateUserRequest(
all_users=True,
user_updates=UpdateUserRequestNoUserIDorEmail.model_validate({field: new_limit}),
),
user_api_key_dict=admin,
litellm_changed_by=None,
)
prisma_client.db.litellm_usertable.update_many.assert_awaited_once_with(where={}, data={field: new_limit})
else:
await user_update(
data=UpdateUserRequest(user_id=saved_user.user_id, **{field: new_limit}),
data=UpdateUserRequest.model_validate({"user_id": saved_user.user_id, field: new_limit}),
user_api_key_dict=admin,
)
assert prisma_client.update_data.call_args.kwargs["data"][field] == new_limit
assert prisma_client.update_data_payload is not None
assert prisma_client.update_data_payload[field] == new_limit
await asyncio.sleep(0)
await asyncio.sleep(0)
remote_subscriber: Final = AuthCacheInvalidationSubscriber(
redis_cache=_FakeRedisCache(),
redis_cache=cast("RedisCache", _FakeRedisCache()),
user_api_key_cache=other_worker_cache,
)
for _, message in published:
@ -2468,13 +2512,13 @@ async def test_user_rate_limit_update_reaches_cached_user_on_every_worker(
handling_user_after: Final = await get_user_object(
user_id=saved_user.user_id,
prisma_client=prisma_client,
prisma_client=prisma_client_for_auth,
user_api_key_cache=handling_worker_cache,
user_id_upsert=False,
)
other_user_after: Final = await get_user_object(
user_id=saved_user.user_id,
prisma_client=prisma_client,
prisma_client=prisma_client_for_auth,
user_api_key_cache=other_worker_cache,
user_id_upsert=False,
)
@ -3080,38 +3124,51 @@ async def test_user_info_v2_proxy_admin_can_query_any_user(mocker: MockerFixture
"""
from fastapi import Request
from litellm.proxy._types import UserInfoV2Response
from litellm.proxy._types import LiteLLM_UserTable, UserInfoV2Response
from litellm.proxy.management_endpoints.internal_user_endpoints import user_info_v2
mock_prisma_client: Final = mocker.MagicMock()
mock_user_row: Final = LiteLLM_UserTable(
user_id="target-user-123",
user_email="target@example.com",
user_alias="Target User",
user_role="internal_user",
spend=42.5,
max_budget=100.0,
tpm_limit=100000,
rpm_limit=1000,
models=["gpt-4"],
budget_duration="30d",
budget_reset_at=None,
metadata={"team": "engineering"},
created_at=datetime(2024, 1, 1, tzinfo=timezone.utc),
updated_at=datetime(2024, 6, 1, tzinfo=timezone.utc),
sso_user_id="sso-abc",
teams=["team-1", "team-2"],
)
mock_user_row: Final = mocker.MagicMock()
mock_user_row.model_dump.return_value = {
"user_id": "target-user-123",
"user_email": "target@example.com",
"user_alias": "Target User",
"user_role": "internal_user",
"spend": 42.5,
"max_budget": 100.0,
"tpm_limit": 100000,
"rpm_limit": 1000,
"models": ["gpt-4"],
"budget_duration": "30d",
"budget_reset_at": None,
"metadata": {"team": "engineering"},
"created_at": datetime(2024, 1, 1, tzinfo=timezone.utc),
"updated_at": datetime(2024, 6, 1, tzinfo=timezone.utc),
"sso_user_id": "sso-abc",
"teams": ["team-1", "team-2"],
}
class _UserTable:
def __init__(self, find_unique: AsyncMock) -> None:
self.find_unique = find_unique
async def mock_find_unique(*_args: object, **kwargs: object) -> MagicMock | None:
class _Database:
def __init__(self, litellm_usertable: _UserTable) -> None:
self.litellm_usertable = litellm_usertable
class _PrismaClient:
def __init__(self, db: _Database) -> None:
self.db = db
async def mock_find_unique(*_args: object, **kwargs: object) -> LiteLLM_UserTable | None:
where: Final = kwargs.get("where")
if isinstance(where, Mapping) and where.get("user_id") == "target-user-123":
return mock_user_row
return None
mock_prisma_client.db.litellm_usertable.find_unique = mocker.AsyncMock(side_effect=mock_find_unique)
mock_prisma_client: Final = _PrismaClient(
db=_Database(
litellm_usertable=_UserTable(find_unique=mocker.AsyncMock(side_effect=mock_find_unique))
)
)
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)