diff --git a/tests/unit/proxy/management_endpoints/test_auto_router_usage.py b/tests/unit/proxy/management_endpoints/test_auto_router_usage.py index dcb7d864fc5..2e9e359450a 100644 --- a/tests/unit/proxy/management_endpoints/test_auto_router_usage.py +++ b/tests/unit/proxy/management_endpoints/test_auto_router_usage.py @@ -1,41 +1,47 @@ -from datetime import date, datetime +from datetime import date +from types import SimpleNamespace from typing import Final -from unittest.mock import AsyncMock +from unittest.mock import ANY, AsyncMock import pytest from fastapi import HTTPException from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.db.prisma_client import PrismaWrapper from litellm.proxy.management_endpoints.auto_router_usage import get_auto_router_usage +from litellm.proxy.utils import PrismaClient class Database: def __init__(self) -> None: - self.get_auto_router_usage = AsyncMock(return_value=()) + self.query_raw = AsyncMock(return_value=()) + self.client = object.__new__(PrismaClient) + self.client.db = PrismaWrapper(original_prisma=SimpleNamespace(query_raw=self.query_raw)) @pytest.mark.asyncio @pytest.mark.parametrize("role", [LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY]) -async def test_admin_filters_reach_prisma_client(role: LitellmUserRoles) -> None: +async def test_admin_filters_reach_query_as_parameters(role: LitellmUserRoles) -> None: db = Database() await get_auto_router_usage( date(2026, 1, 1), date(2026, 1, 2), UserAPIKeyAuth(user_role=role), - db, + db.client, router_name="router'quoted", router_type="complexity", user_id="owner", api_key="key-hash", ) - db.get_auto_router_usage.assert_awaited_once_with( - start_time=datetime(2026, 1, 1), - end_time=datetime(2026, 1, 3), - user_id="owner", - api_key="key-hash", - router_name="router'quoted", - router_type="complexity", - destination_model=None, + db.query_raw.assert_awaited_once_with( + ANY, + "2026-01-01T00:00:00", + "2026-01-03T00:00:00", + "owner", + "key-hash", + "router'quoted", + "complexity", + None, ) @@ -46,17 +52,18 @@ async def test_non_admin_is_scoped_to_own_user_when_filter_omitted() -> None: date(2026, 1, 1), date(2026, 1, 1), UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="own-user"), - db, + db.client, destination_model="model-a", ) - db.get_auto_router_usage.assert_awaited_once_with( - start_time=datetime(2026, 1, 1), - end_time=datetime(2026, 1, 2), - user_id="own-user", - api_key=None, - router_name=None, - router_type=None, - destination_model="model-a", + db.query_raw.assert_awaited_once_with( + ANY, + "2026-01-01T00:00:00", + "2026-01-02T00:00:00", + "own-user", + None, + None, + None, + "model-a", ) @@ -69,12 +76,12 @@ async def test_non_admin_cannot_read_other_users_or_unbound_service_account(call date(2026, 1, 1), date(2026, 1, 1), UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id=caller), - db, + db.client, destination_model="model-a", user_id="another-user", ) assert error.value.status_code == 403 - db.get_auto_router_usage.assert_not_called() + db.query_raw.assert_not_called() @pytest.mark.asyncio @@ -86,12 +93,12 @@ async def test_query_requires_one_specific_model_or_router(model: str | None, ro date(2026, 1, 1), date(2026, 1, 1), UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), - db, + db.client, destination_model=model, router_name=router_name, ) assert error.value.status_code == 400 - db.get_auto_router_usage.assert_not_called() + db.query_raw.assert_not_called() @pytest.mark.asyncio @@ -100,21 +107,29 @@ async def test_invalid_or_overlong_ranges_never_query_spend_logs(end: date) -> N db: Final = Database() with pytest.raises(HTTPException) as error: await get_auto_router_usage( - date(2026, 1, 1), end, UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), db, destination_model="fast" + date(2026, 1, 1), end, UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), db.client, destination_model="fast" ) assert error.value.status_code == 400 - db.get_auto_router_usage.assert_not_called() + db.query_raw.assert_not_called() @pytest.mark.asyncio async def test_maximum_range_includes_its_last_day() -> None: db: Final = Database() await get_auto_router_usage( - date(2026, 1, 1), date(2026, 4, 3), UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), db, + date(2026, 1, 1), date(2026, 4, 3), UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), db.client, destination_model="fast", ) - assert db.get_auto_router_usage.call_args.kwargs["start_time"] == datetime(2026, 1, 1) - assert db.get_auto_router_usage.call_args.kwargs["end_time"] == datetime(2026, 4, 4) + db.query_raw.assert_awaited_once_with( + ANY, + "2026-01-01T00:00:00", + "2026-04-04T00:00:00", + None, + None, + None, + None, + "fast", + ) @pytest.mark.asyncio @@ -122,8 +137,8 @@ async def test_router_type_cannot_filter_a_destination_model() -> None: db: Final = Database() with pytest.raises(HTTPException) as error: await get_auto_router_usage( - date(2026, 1, 1), date(2026, 1, 1), UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), db, + date(2026, 1, 1), date(2026, 1, 1), UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), db.client, destination_model="fast", router_type="complexity", ) assert error.value.status_code == 400 - db.get_auto_router_usage.assert_not_called() + db.query_raw.assert_not_called()