From 00f6210f7fe6de21c441bd3bc910036d346b2ec6 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 24 Sep 2026 18:20:34 -0700 Subject: [PATCH] refactor(ptu): keep Azure detection under llms/azure and check share teams with one query --- litellm/llms/azure/ptu_capacity.py | 19 +++++++++++++++++ .../model_management_endpoints.py | 18 +++++++++++++--- litellm/router_utils/ptu_shares.py | 20 +++--------------- .../test_ptu_model_settings.py | 21 ++++++++++++------- 4 files changed, 50 insertions(+), 28 deletions(-) diff --git a/litellm/llms/azure/ptu_capacity.py b/litellm/llms/azure/ptu_capacity.py index c68429dc7b7..ce50e885b66 100644 --- a/litellm/llms/azure/ptu_capacity.py +++ b/litellm/llms/azure/ptu_capacity.py @@ -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: diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index ab33c35909a..41df67557e7 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -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"}) diff --git a/litellm/router_utils/ptu_shares.py b/litellm/router_utils/ptu_shares.py index 3bc9d33349c..11b61676ce5 100644 --- a/litellm/router_utils/ptu_shares.py +++ b/litellm/router_utils/ptu_shares.py @@ -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" 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 cc1ac071207..b81ae11ac5d 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,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()