fix(ptu): refuse ptu_shares naming a team that does not exist

This commit is contained in:
mateo-berri 2026-09-24 17:53:51 -07:00
parent 4ad0db30b2
commit 6a24e7f920
2 changed files with 101 additions and 3 deletions

View file

@ -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:

View file

@ -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(