From f552e86fea57b8b75494958e2b92d716d0775895 Mon Sep 17 00:00:00 2001 From: mrinal Date: Fri, 2 Oct 2026 01:08:46 +0000 Subject: [PATCH] test(proxy): type user endpoint test doubles Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../test_internal_user_endpoints.py | 155 ++++++++++++------ 1 file changed, 106 insertions(+), 49 deletions(-) diff --git a/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py index 728e156c3e6..6ab8c431573 100644 --- a/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py @@ -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)