mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
fix(proxy): chunk the ptu_shares team lookup under the IN-list bound
This commit is contained in:
parent
48c0be34aa
commit
d1f698fdd8
2 changed files with 16 additions and 9 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue