From d1f698fdd8cb82f4319c4e438746d647602a8a04 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 16:34:00 -0700 Subject: [PATCH] fix(proxy): chunk the ptu_shares team lookup under the IN-list bound --- .../model_management_endpoints.py | 11 +++-------- .../test_ptu_model_settings.py | 14 +++++++++++++- 2 files changed, 16 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index b16ef81e6c2..0449ccccee2 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -24,7 +24,6 @@ from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, TypeAlias from fastapi import APIRouter, Depends, Header, HTTPException, Request, status from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError, field_validator -from typing_extensions import ReadOnly, TypedDict import litellm from litellm._logging import verbose_proxy_logger @@ -99,6 +98,7 @@ from litellm.proxy.spend_tracking.ptu_feature_flag import ( is_ptu_cost_attribution_enabled, ) from litellm.proxy.utils import PrismaClient, ProxyLogging +from litellm.repositories.chunked_in import find_many_in from litellm.repositories.credentials_repository import CredentialsRepository from litellm.repositories.model_repository import ModelRepository from litellm.repositories.prisma_protocols import TableActions @@ -247,14 +247,10 @@ class _TeamRow(Protocol): def model_dump(self) -> Mapping[str, object]: ... -class _TeamIdsWhere(TypedDict): - team_id: ReadOnly[Mapping[Literal["in"], Sequence[str]]] - - class _TeamLookupTable(Protocol): def find_unique(self, *, where: Mapping[str, object]) -> Awaitable[_TeamRow | None]: ... - def find_many(self, *, where: Mapping[str, object]) -> Awaitable[Sequence[_TeamRow]]: ... + async def find_many(self, *, where: Mapping[str, object]) -> Sequence[_TeamRow]: ... class _TeamTable(_TeamLookupTable, Protocol): @@ -798,8 +794,7 @@ async def _raise_if_ptu_share_teams_missing( shares: Final = parsed_ptu_shares(model_info.get("ptu_shares")) if shares is None: return - where: Final[_TeamIdsWhere] = {"team_id": {"in": tuple(shares)}} - rows: Final = await team_table().find_many(where=where) + rows: Final = await find_many_in(team_table(), "team_id", shares.keys()) found: Final = frozenset(row.team_id for row in rows) missing: Final = tuple(team_id for team_id in shares if team_id not in found) if not missing: diff --git a/tests/test_litellm/proxy/management_endpoints/test_ptu_model_settings.py b/tests/test_litellm/proxy/management_endpoints/test_ptu_model_settings.py index b81ae11ac5d..ce2e865d69d 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ptu_model_settings.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ptu_model_settings.py @@ -38,6 +38,7 @@ from litellm.proxy.management_endpoints.model_management_endpoints import ( update_db_model, ) from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR +from litellm.repositories.chunked_in import IN_LIST_CHUNK_SIZE from litellm.router import Router from litellm.types.router import ( SPECIAL_MODEL_INFO_PARAMS, @@ -1337,6 +1338,7 @@ class _TeamLookup: def __init__(self, existing: frozenset[str]) -> None: self.existing: Final = existing self.looked_up: tuple[str, ...] = () + self.batch_sizes: tuple[int, ...] = () async def find_unique(self, *, where: Mapping[str, object]) -> LiteLLM_TeamTable | None: raise AssertionError(f"one lookup per team is what the review asked to avoid: {where}") @@ -1346,6 +1348,7 @@ class _TeamLookup: assert isinstance(team_filter, Mapping) requested: Final = tuple(str(team_id) for team_id in cast(Sequence[object], team_filter["in"])) self.looked_up = (*self.looked_up, *requested) + self.batch_sizes = (*self.batch_sizes, len(requested)) return tuple(LiteLLM_TeamTable(team_id=team_id) for team_id in requested if team_id in self.existing) @@ -1375,6 +1378,15 @@ async def test_share_team_check_accepts_shares_naming_existing_teams(): assert lookup.looked_up == ("team-a", "team-b") +@pytest.mark.asyncio +async def test_share_team_check_splits_a_share_list_longer_than_one_in_list_chunk(): + team_ids: Final = tuple(f"team-{index}" for index in range(IN_LIST_CHUNK_SIZE + 1)) + lookup: Final = _TeamLookup(frozenset(team_ids)) + await _raise_if_ptu_share_teams_missing(_shared_model_info(dict.fromkeys(team_ids, 1)), lambda: lookup) + assert lookup.looked_up == team_ids + assert max(lookup.batch_sizes) <= IN_LIST_CHUNK_SIZE + + @pytest.mark.asyncio async def test_share_team_check_leaves_a_team_id_holder_to_the_team_model_check(): lookup: Final = _TeamLookup(frozenset()) @@ -1406,7 +1418,7 @@ async def test_model_new_refuses_shares_naming_a_team_that_does_not_exist(monkey await add_new_model(model_params=shared, user_api_key_dict=admin) assert exc.value.code == "400" - prisma_client.db.litellm_teamtable.find_many.assert_awaited_once_with(where={"team_id": {"in": ("ghost-team",)}}) + prisma_client.db.litellm_teamtable.find_many.assert_awaited_once_with(where={"team_id": {"in": ["ghost-team"]}}) add_model_to_db.assert_not_called() add_team_model_to_db.assert_not_called()