diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index ae294871afc..ab33c35909a 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -17,6 +17,7 @@ from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequen from contextlib import AbstractAsyncContextManager, asynccontextmanager, suppress from dataclasses import dataclass from fnmatch import fnmatchcase +from functools import partial from json import JSONDecodeError from types import MappingProxyType from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, TypeAlias, TypeVar, cast, runtime_checkable @@ -37,6 +38,7 @@ from litellm.litellm_core_utils.ptu_pricing import ( PTU_ZEROED_PRICING_FIELDS, PTU_ZEROED_TABLE_FIELDS, SEARCH_CONTEXT_SIZES, + parsed_ptu_shares, ptu_config_error, ) from litellm.proxy._types import ( @@ -779,6 +781,21 @@ def _validate_ptu_model_info(model_info: Mapping[str, object]) -> None: raise HTTPException(status_code=400, detail=error) +async def _raise_if_ptu_share_teams_missing( + model_info: Mapping[str, object], team_table: Callable[[], _TeamLookupTable] +) -> None: + """Hold every team named in ``ptu_shares`` to the existence check ``team_id`` already gets.""" + shares: Final = parsed_ptu_shares(model_info.get("ptu_shares")) + if shares is None: + return + table: Final = team_table() + rows: Final = await asyncio.gather(*(table.find_unique(where={"team_id": team_id}) for team_id in shares)) + missing: Final = tuple(team_id for team_id, row in zip(shares, rows, strict=True) if row is None) + if not missing: + return + raise HTTPException(status_code=400, detail={"error": f"Team id={', '.join(missing)} does not exist in db"}) + + # The mirrored per-token pricing fields plus the remaining rates the public cost map or a # provider default would otherwise supply (the cache back-fills, the Maps grounding rate). An # unset field falls back to those sources, so a field left out here is one a PTU deployment @@ -1614,8 +1631,10 @@ async def _update_team_model_in_db( # raising the rate on a configured model carries no ptu_effective_from, which the # stored row supplies. if patch_data.model_info is not None: - _raise_if_ptu_cost_attribution_disabled(patch_data.model_info.model_dump(exclude_none=True)) + incoming_model_info: Final = patch_data.model_info.model_dump(exclude_none=True) + _raise_if_ptu_cost_attribution_disabled(incoming_model_info) _validate_ptu_model_info(_merged_ptu_model_info(db_model=db_model, patch_data=patch_data)) + await _raise_if_ptu_share_teams_missing(incoming_model_info, partial(_repo_team_table, prisma_client)) _raise_if_ptu_deployment_is_priced( model_info=_merged_ptu_model_info(db_model=db_model, patch_data=patch_data), supplied=( @@ -2456,6 +2475,7 @@ async def add_new_model( incoming_model_info: Final = model_params.model_info.model_dump(exclude_none=True) _raise_if_ptu_cost_attribution_disabled(incoming_model_info) _validate_ptu_model_info(incoming_model_info) + await _raise_if_ptu_share_teams_missing(incoming_model_info, partial(_repo_team_table, prisma_client)) priced_model_params: Final = _ptu_priced_deployment(model_params) if store_model_in_db is True: 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 8ed3df44216..cc1ac071207 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 @@ -2,7 +2,9 @@ import datetime import json +from collections.abc import Mapping from contextlib import ExitStack +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch from unittest.mock import patch as patch_ctx @@ -14,7 +16,9 @@ from litellm.litellm_core_utils.ptu_pricing import ptu_terms from litellm.llms.gemini.cost_calculator import cost_per_web_search_request from litellm.proxy._types import ( LiteLLM_ProxyModelTable, + LiteLLM_TeamTable, LitellmUserRoles, + ProxyException, ReconcileOutcome, UserAPIKeyAuth, ) @@ -27,6 +31,7 @@ from litellm.proxy.management_endpoints.model_management_endpoints import ( _ptu_priced_deployment, _ptu_zeroed_pricing, _raise_if_ptu_cost_attribution_disabled, + _raise_if_ptu_share_teams_missing, _update_team_model_in_db, _validate_ptu_model_info, add_new_model, @@ -651,7 +656,7 @@ class TestAddNewModelPtuGate: monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False) @staticmethod - def _patched_proxy(model_id: str): + def _patched_proxy(model_id: str, prisma_client: MagicMock | None = None): """Patch everything /model/new touches except the PTU gate, and hand back the DB writers.""" db_row = LiteLLM_ProxyModelTable( model_id=model_id, @@ -678,7 +683,7 @@ class TestAddNewModelPtuGate: proxy_server = "litellm.proxy.proxy_server" endpoints = "litellm.proxy.management_endpoints.model_management_endpoints" return (add_model_to_db, add_team_model_to_db), [ - patch(f"{proxy_server}.prisma_client", MagicMock()), + patch(f"{proxy_server}.prisma_client", prisma_client if prisma_client is not None else MagicMock()), patch(f"{proxy_server}.store_model_in_db", True), patch(f"{proxy_server}.proxy_config", mock_proxy_config), patch(f"{proxy_server}.proxy_logging_obj", MagicMock()), @@ -1328,6 +1333,79 @@ def test_validate_helper_refuses_shares_that_do_not_add_up_to_the_count(): assert "4 of 5 allocated" in exc.value.detail +class _TeamLookup: + def __init__(self, existing: frozenset[str]) -> None: + self.existing: Final = existing + self.looked_up: tuple[str, ...] = () + + async def find_unique(self, *, where: Mapping[str, object]) -> LiteLLM_TeamTable | None: + team_id: Final = str(where["team_id"]) + self.looked_up = (*self.looked_up, team_id) + return LiteLLM_TeamTable(team_id=team_id) if team_id in self.existing else None + + +def _shared_model_info(shares: Mapping[str, int]) -> Mapping[str, object]: + return { + "ptu_count": sum(shares.values()), + "cost_per_ptu_per_hour": 2.0, + "ptu_effective_from": _SHARED_START, + "ptu_shares": dict(shares), + } + + +@pytest.mark.asyncio +async def test_share_team_check_refuses_a_team_that_does_not_exist(): + with pytest.raises(HTTPException) as exc: + await _raise_if_ptu_share_teams_missing( + _shared_model_info({"team-a": 3, "ghost-team": 2}), lambda: _TeamLookup(frozenset({"team-a"})) + ) + assert exc.value.status_code == 400 + assert exc.value.detail == {"error": "Team id=ghost-team does not exist in db"} + + +@pytest.mark.asyncio +async def test_share_team_check_accepts_shares_naming_existing_teams(): + lookup: Final = _TeamLookup(frozenset({"team-a", "team-b"})) + await _raise_if_ptu_share_teams_missing(_shared_model_info({"team-a": 3, "team-b": 2}), lambda: lookup) + assert sorted(lookup.looked_up) == ["team-a", "team-b"] + + +@pytest.mark.asyncio +async def test_share_team_check_leaves_a_team_id_holder_to_the_team_model_check(): + lookup: Final = _TeamLookup(frozenset()) + await _raise_if_ptu_share_teams_missing( + {"team_id": "team-a", "ptu_count": 5, "cost_per_ptu_per_hour": 2.0, "ptu_effective_from": _SHARED_START}, + lambda: lookup, + ) + assert lookup.looked_up == () + + +@pytest.mark.asyncio +async def test_model_new_refuses_shares_naming_a_team_that_does_not_exist(monkeypatch): + monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true") + prisma_client = MagicMock() + prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None) + (add_model_to_db, add_team_model_to_db), patches = TestAddNewModelPtuGate._patched_proxy( + "ptu-shared-model", prisma_client=prisma_client + ) + admin = UserAPIKeyAuth(user_id="test-admin", user_role=LitellmUserRoles.PROXY_ADMIN) + base = TestAddNewModelPtuGate._ptu_deployment("ptu-shared-model") + shared = base.model_copy( + update={"model_info": base.model_info.model_copy(update={"team_id": None, "ptu_shares": {"ghost-team": 15}})} + ) + + with ExitStack() as stack: + for active_patch in patches: + stack.enter_context(active_patch) + with pytest.raises(ProxyException, match="Team id=ghost-team does not exist in db") as exc: + await add_new_model(model_params=shared, user_api_key_dict=admin) + + assert exc.value.code == "400" + prisma_client.db.litellm_teamtable.find_unique.assert_awaited_once_with(where={"team_id": "ghost-team"}) + add_model_to_db.assert_not_called() + add_team_model_to_db.assert_not_called() + + def test_validate_helper_refuses_a_team_id_beside_shares(): with pytest.raises(HTTPException) as exc: _validate_ptu_model_info(