mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
refactor(ptu): keep Azure detection under llms/azure and check share teams with one query
This commit is contained in:
parent
6a24e7f920
commit
00f6210f7f
4 changed files with 50 additions and 28 deletions
|
|
@ -17,6 +17,8 @@ from dataclasses import dataclass
|
|||
from types import MappingProxyType
|
||||
from typing import Final, Protocol
|
||||
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
|
||||
class NormalizedTokenWeights(Protocol):
|
||||
@property
|
||||
|
|
@ -36,6 +38,9 @@ class PTUCapacity:
|
|||
def normalized_tokens_per_ptu_hour(self) -> int:
|
||||
return self.input_tpm_per_ptu * 60
|
||||
|
||||
def input_tpm_for(self, ptus: int) -> int:
|
||||
return ptus * self.input_tpm_per_ptu
|
||||
|
||||
|
||||
AZURE_PTU_CAPACITY: Final[Mapping[str, PTUCapacity]] = MappingProxyType(
|
||||
{
|
||||
|
|
@ -68,6 +73,7 @@ AZURE_PTU_CAPACITY: Final[Mapping[str, PTUCapacity]] = MappingProxyType(
|
|||
)
|
||||
|
||||
_VERSION_SUFFIX: Final = re.compile(r"-\d{4}-\d{2}-\d{2}$")
|
||||
_AZURE_PROVIDERS: Final = frozenset({LlmProviders.AZURE.value, LlmProviders.AZURE_AI.value})
|
||||
|
||||
|
||||
def azure_ptu_capacity(model: str) -> PTUCapacity | None:
|
||||
|
|
@ -97,6 +103,19 @@ def deployment_ptu_capacity(deployment: Mapping[str, object]) -> PTUCapacity | N
|
|||
return next((capacity for capacity in map(azure_ptu_capacity, candidates) if capacity is not None), None)
|
||||
|
||||
|
||||
def is_azure_deployment(deployment: Mapping[str, object]) -> bool:
|
||||
"""Whether ``litellm_params`` route this deployment to Azure OpenAI or Azure AI, by
|
||||
``custom_llm_provider`` first and the ``model`` prefix otherwise."""
|
||||
litellm_params: Final = deployment.get("litellm_params")
|
||||
if not isinstance(litellm_params, Mapping):
|
||||
return False
|
||||
provider: Final = litellm_params.get("custom_llm_provider")
|
||||
if isinstance(provider, str):
|
||||
return provider in _AZURE_PROVIDERS
|
||||
model: Final = litellm_params.get("model")
|
||||
return isinstance(model, str) and model.partition("/")[0] in _AZURE_PROVIDERS
|
||||
|
||||
|
||||
def normalized_tokens(
|
||||
weights: NormalizedTokenWeights, *, prompt_tokens: int, completion_tokens: int, cache_read_tokens: int = 0
|
||||
) -> float:
|
||||
|
|
|
|||
|
|
@ -24,6 +24,7 @@ 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
|
||||
|
|
@ -134,6 +135,7 @@ from litellm.types.proxy.management_endpoints.model_management_endpoints import
|
|||
AutoRouterClassifierDefaultPromptResponse,
|
||||
UpdateUsefulLinksRequest,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.team_endpoints import TeamIdSearchFilter
|
||||
from litellm.types.router import (
|
||||
SPECIAL_MODEL_INFO_PARAMS,
|
||||
Deployment,
|
||||
|
|
@ -237,15 +239,24 @@ class _ExistingModelRow(Protocol):
|
|||
|
||||
|
||||
class _TeamRow(Protocol):
|
||||
@property
|
||||
def team_id(self) -> str: ...
|
||||
|
||||
@property
|
||||
def models(self) -> Sequence[str]: ...
|
||||
|
||||
def model_dump(self) -> Mapping[str, object]: ...
|
||||
|
||||
|
||||
class _TeamIdsWhere(TypedDict):
|
||||
team_id: ReadOnly[TeamIdSearchFilter]
|
||||
|
||||
|
||||
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]]: ...
|
||||
|
||||
|
||||
class _TeamTable(_TeamLookupTable, Protocol):
|
||||
def update(
|
||||
|
|
@ -788,9 +799,10 @@ async def _raise_if_ptu_share_teams_missing(
|
|||
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)
|
||||
where: Final[_TeamIdsWhere] = {"team_id": {"in": tuple(shares)}}
|
||||
rows: Final = await team_table().find_many(where=where)
|
||||
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:
|
||||
return
|
||||
raise HTTPException(status_code=400, detail={"error": f"Team id={', '.join(missing)} does not exist in db"})
|
||||
|
|
|
|||
|
|
@ -11,10 +11,7 @@ from dataclasses import dataclass
|
|||
from typing import Final, Generic, TypeVar
|
||||
|
||||
from litellm.litellm_core_utils.ptu_pricing import parsed_ptu_shares, ptu_terms
|
||||
from litellm.llms.azure.ptu_capacity import PTUCapacity, deployment_ptu_capacity
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
_AZURE_PROVIDERS: Final = frozenset({LlmProviders.AZURE.value, LlmProviders.AZURE_AI.value})
|
||||
from litellm.llms.azure.ptu_capacity import PTUCapacity, deployment_ptu_capacity, is_azure_deployment
|
||||
|
||||
_DeploymentT = TypeVar("_DeploymentT", bound=Mapping[str, object])
|
||||
|
||||
|
|
@ -76,7 +73,7 @@ def team_ptu_ceiling(deployments: Sequence[Mapping[str, object]], team_id: str)
|
|||
if not priced:
|
||||
return None
|
||||
return PTUTeamCeiling(
|
||||
tpm_limit=sum(share * capacity.input_tpm_per_ptu for share, capacity in priced),
|
||||
tpm_limit=sum(capacity.input_tpm_for(share) for share, capacity in priced),
|
||||
output_to_input_ratio=max(capacity.output_to_input_ratio for _, capacity in priced),
|
||||
cached_input_ratio=max(capacity.cached_input_ratio for _, capacity in priced),
|
||||
)
|
||||
|
|
@ -112,17 +109,6 @@ def model_group_ptu_capacity(deployments: Sequence[Mapping[str, object]]) -> PTU
|
|||
)
|
||||
|
||||
|
||||
def _is_azure_deployment(deployment: Mapping[str, object]) -> bool:
|
||||
litellm_params: Final = deployment.get("litellm_params")
|
||||
if not isinstance(litellm_params, Mapping):
|
||||
return False
|
||||
provider: Final = litellm_params.get("custom_llm_provider")
|
||||
if isinstance(provider, str):
|
||||
return provider in _AZURE_PROVIDERS
|
||||
model: Final = litellm_params.get("model")
|
||||
return isinstance(model, str) and model.partition("/")[0] in _AZURE_PROVIDERS
|
||||
|
||||
|
||||
def ptu_capacity_warning(model_name: str, deployment: Mapping[str, object]) -> str | None:
|
||||
"""Why this reserved deployment's tokens cannot be converted to PTUs, else None.
|
||||
|
||||
|
|
@ -141,7 +127,7 @@ def ptu_capacity_warning(model_name: str, deployment: Mapping[str, object]) -> s
|
|||
f"PTU deployment '{model_name}' has no Azure sizing row for its model, so its PTU shares set no "
|
||||
"team ceiling and its usage reports no PTU hours; set model_info.base_model to the Azure model name"
|
||||
)
|
||||
if _is_azure_deployment(deployment):
|
||||
if is_azure_deployment(deployment):
|
||||
return (
|
||||
f"PTU deployment '{model_name}' has no Azure sizing row for its model, so its usage reports no "
|
||||
"PTU hours; set model_info.base_model to the Azure model name"
|
||||
|
|
|
|||
|
|
@ -2,9 +2,9 @@
|
|||
|
||||
import datetime
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Mapping, Sequence
|
||||
from contextlib import ExitStack
|
||||
from typing import Final
|
||||
from typing import Final, cast
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
from unittest.mock import patch as patch_ctx
|
||||
|
||||
|
|
@ -1339,9 +1339,14 @@ class _TeamLookup:
|
|||
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
|
||||
raise AssertionError(f"one lookup per team is what the review asked to avoid: {where}")
|
||||
|
||||
async def find_many(self, *, where: Mapping[str, object]) -> Sequence[LiteLLM_TeamTable]:
|
||||
team_filter: Final = where["team_id"]
|
||||
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)
|
||||
return tuple(LiteLLM_TeamTable(team_id=team_id) for team_id in requested if team_id in self.existing)
|
||||
|
||||
|
||||
def _shared_model_info(shares: Mapping[str, int]) -> Mapping[str, object]:
|
||||
|
|
@ -1367,7 +1372,7 @@ async def test_share_team_check_refuses_a_team_that_does_not_exist():
|
|||
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"]
|
||||
assert lookup.looked_up == ("team-a", "team-b")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -1384,7 +1389,7 @@ async def test_share_team_check_leaves_a_team_id_holder_to_the_team_model_check(
|
|||
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)
|
||||
prisma_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[])
|
||||
(add_model_to_db, add_team_model_to_db), patches = TestAddNewModelPtuGate._patched_proxy(
|
||||
"ptu-shared-model", prisma_client=prisma_client
|
||||
)
|
||||
|
|
@ -1401,7 +1406,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_unique.assert_awaited_once_with(where={"team_id": "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