mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
fix(ptu): refuse ptu_shares naming a team that does not exist
This commit is contained in:
parent
4ad0db30b2
commit
6a24e7f920
2 changed files with 101 additions and 3 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue