mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
88195a49ce
commit
f552e86fea
1 changed files with 106 additions and 49 deletions
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue