refactor(ptu): keep Azure detection under llms/azure and check share teams with one query

This commit is contained in:
mateo-berri 2026-09-24 18:20:34 -07:00
parent 6a24e7f920
commit 00f6210f7f
4 changed files with 50 additions and 28 deletions

View file

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

View file

@ -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"})

View file

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

View file

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