From 5013833136bde13f71be22e8ed8478bbe37c6ba0 Mon Sep 17 00:00:00 2001 From: Tomoya Tabuchi Date: Tue, 9 Jun 2026 16:12:46 +0900 Subject: [PATCH] feat(team_endpoints): add tests for limitting key count in /team/info response --- .../test_team_endpoints.py | 33 +++++++++++++++++++ .../test_prisma_client_get_data.py | 31 +++++++++++++++-- 2 files changed, 62 insertions(+), 2 deletions(-) diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index d580f1f7703..ffb740a9178 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -8306,6 +8306,8 @@ async def test_new_team_encrypts_callback_vars( assert cv["langfuse_secret_key"] != "sk-real" recovered = decrypt_callback_vars(metadata)["logging"][0]["callback_vars"] assert recovered["langfuse_secret_key"] == "sk-real" + + def _non_admin_auth(): return UserAPIKeyAuth( user_id="u-team-admin", user_role=LitellmUserRoles.INTERNAL_USER @@ -8404,3 +8406,34 @@ async def test_update_team_blocks_non_admin_passthrough_routes(mock_db_client): ) assert str(exc.value.code) == "403" assert "allowed_passthrough_routes" in str(exc.value.message) + + +@pytest.mark.asyncio +async def test_team_info_forwards_key_limit_to_get_data(): + """/team/info must thread its ``key_limit`` query param into the key + lookup so the database caps how many keys are returned for the team. + """ + from fastapi import Request + + from litellm.proxy.management_endpoints import team_endpoints + + mock_prisma = MagicMock() + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( + return_value=LiteLLM_TeamTable(team_id="team-1") + ) + mock_prisma.get_data = AsyncMock(return_value=[]) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch.object( + team_endpoints, "get_all_team_memberships", AsyncMock(return_value=[]) + ), + ): + await team_endpoints.team_info( + http_request=MagicMock(spec=Request), + team_id="team-1", + key_limit=7, + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), + ) + + assert mock_prisma.get_data.await_args.kwargs["limit"] == 7 diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_get_data.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_get_data.py index 7e7e98d1360..4074750788b 100644 --- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_get_data.py +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_get_data.py @@ -234,7 +234,9 @@ async def test_query_first_with_cached_plan_fallback_retries_on_cached_plan_erro async def test_query_first_with_cached_plan_fallback_reraises_non_plan_errors( prisma_client: PrismaClient, ) -> None: - prisma_client.db.query_first = AsyncMock(side_effect=RuntimeError("totally unrelated")) + prisma_client.db.query_first = AsyncMock( + side_effect=RuntimeError("totally unrelated") + ) with pytest.raises(RuntimeError, match="totally unrelated"): await prisma_client._query_first_with_cached_plan_fallback("SELECT 1") @@ -351,7 +353,9 @@ async def test_get_data_token_find_unique_returns_record( async def test_get_data_token_find_unique_missing_token_raises_401( prisma_client: PrismaClient, ) -> None: - prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=None) + prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=None + ) with pytest.raises(HTTPException) as excinfo: await prisma_client.get_data(token="sk-missing", table_name="key") err = excinfo.value @@ -398,3 +402,26 @@ async def test_get_data_logs_and_raises_on_db_error( ) with pytest.raises(RuntimeError, match="network split"): await prisma_client.get_data(token="sk-broken", table_name="key") + + +@pytest.mark.asyncio +@pytest.mark.parametrize("limit", [5, None]) +async def test_get_data_team_keys_forward_limit_as_take( + prisma_client: PrismaClient, limit: Any +) -> None: + """The /team/info ``key_limit`` must reach Prisma as ``take`` so the + database caps how many of a team's keys come back. + ``limit=None`` leaves ``take`` unset so every key is returned. + """ + prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + await prisma_client.get_data( + team_id="team-1", + table_name="key", + query_type="find_all", + limit=limit, + ) + assert prisma_client.db.litellm_verificationtoken.find_many.await_args.kwargs == { + "take": limit, + "where": {"team_id": "team-1"}, + "include": {"litellm_budget_table": True}, + }