fix(proxy): weight cached input in PTU ceilings and size team-scoped deployments by public name

The PTU flag is read from the process environment on every check instead of
through the secret-manager path, so a proxy with a hosted secret manager does
not pay a round trip per request. A team's ceiling now carries the model's
cached-input ratio and settlement charges cache reads at that ratio, matching
how Azure sizes a PTU. A deployment registered for one team is found by its
team_public_model_name for PTU-hours and the ceiling, since its model_name is
rewritten to a unique internal name. ptu_shares are StrictInt on ModelInfo and
parsed_ptu_shares refuses booleans, strings, and non-string team ids, a
fractional ptu_count reserves nothing, and models_processed counts distinct
deployments rather than holdings.
This commit is contained in:
mateo-berri 2026-09-24 14:44:47 -07:00
parent 95e3caefb0
commit 304e9d2ca0
15 changed files with 448 additions and 95 deletions

View file

@ -15,7 +15,15 @@ import re
from collections.abc import Mapping
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final
from typing import Final, Protocol
class NormalizedTokenWeights(Protocol):
@property
def output_to_input_ratio(self) -> float: ...
@property
def cached_input_ratio(self) -> float: ...
@dataclass(frozen=True, slots=True)
@ -90,13 +98,13 @@ def deployment_ptu_capacity(deployment: Mapping[str, object]) -> PTUCapacity | N
def normalized_tokens(
capacity: PTUCapacity, *, prompt_tokens: int, completion_tokens: int, cache_read_tokens: int = 0
weights: NormalizedTokenWeights, *, prompt_tokens: int, completion_tokens: int, cache_read_tokens: int = 0
) -> float:
"""Azure's normalized token count for one request: uncached input in full, cached input
at the model's cached ratio, output weighted by the output-to-input ratio."""
cached: Final = min(max(cache_read_tokens, 0), max(prompt_tokens, 0))
uncached: Final = max(prompt_tokens, 0) - cached
return uncached + capacity.cached_input_ratio * cached + capacity.output_to_input_ratio * max(completion_tokens, 0)
return uncached + weights.cached_input_ratio * cached + weights.output_to_input_ratio * max(completion_tokens, 0)
def ptu_hours(capacity: PTUCapacity, normalized: float) -> float:

View file

@ -6,13 +6,14 @@ together because they have to agree: a deployment the rollup declines to charge
router prices at zero serves its traffic for free.
"""
import os
from collections.abc import Mapping
from dataclasses import dataclass
from datetime import date, datetime, time, timezone
from types import MappingProxyType
from typing import Final
from litellm.secret_managers.main import get_secret_bool
from litellm.secret_managers.main import str_to_bool
from litellm.types.router import ModelInfo
from litellm.types.utils import AzureSpillover, CustomPricingLiteLLMParams, MirroredPricingParams
@ -22,8 +23,12 @@ AZURE_SPILLOVER_FROM_HEADER: Final = "x-ms-spillover-from-deployment"
def is_ptu_cost_attribution_enabled() -> bool:
"""Whether PTU flat-cost attribution is turned on for this process."""
return get_secret_bool(PTU_COST_ATTRIBUTION_ENV_VAR, False) is True
"""Whether PTU flat-cost attribution is turned on for this process.
Read from the environment alone: the router and the rate limiter ask on every request,
and ``get_secret`` would forward each of those reads to a configured secret manager.
"""
return str_to_bool(os.environ.get(PTU_COST_ATTRIBUTION_ENV_VAR)) is True
PTU_ZEROED_PRICING_FIELDS: Final = tuple(f for f in MirroredPricingParams.model_fields if f != "tiered_pricing") + (
@ -148,14 +153,30 @@ def parsed_ptu_shares(raw: object) -> Mapping[str, int] | None:
"""
if not isinstance(raw, Mapping) or not raw:
return None
entries: Final = tuple((str(team_id), share) for team_id, share in raw.items())
if any(
not team_id or isinstance(share, bool) or not isinstance(share, int) or share <= 0 for team_id, share in entries
):
entries: Final = tuple(
(team_id, share)
for team_id, share in raw.items()
if isinstance(team_id, str) and team_id and isinstance(share, int) and not isinstance(share, bool) and share > 0
)
if len(entries) != len(raw):
return None
return MappingProxyType(dict(entries))
def _parsed_ptu_count(model_info: Mapping[str, object]) -> int | None:
"""``ptu_count`` as the whole number of reserved units within bounds, else None."""
raw: Final = model_info.get("ptu_count")
if isinstance(raw, bool) or not isinstance(raw, (int, float, str)):
return None
if isinstance(raw, float) and not raw.is_integer():
return None
try:
count: Final = int(raw)
except (ValueError, OverflowError):
return None
return count if 0 < count <= ModelInfo.MAX_PTU_COUNT else None
def _declared_shares(model_info: Mapping[str, object], ptu_count: int) -> Mapping[str, int] | None:
"""Who holds the capacity: the single ``team_id`` holding all of it, or the ``ptu_shares``
that add up to it, else None."""
@ -229,9 +250,8 @@ def _ptu_holder_error(model_info: Mapping[str, object], model_name: str | None)
shares: Final = parsed_ptu_shares(raw_shares)
if shares is None:
return _named("ptu_shares must map at least one team_id to a positive whole number of PTUs", model_name)
try:
ptu_count: Final = int(str(model_info.get("ptu_count")))
except ValueError:
ptu_count: Final = _parsed_ptu_count(model_info)
if ptu_count is None:
return None
allocated: Final = sum(shares.values())
if allocated != ptu_count:
@ -246,16 +266,13 @@ def ptu_terms(model_info: Mapping[str, object]) -> PTUTerms | None:
present but unparseable bound would read as no bound and widen the window to the whole
day, so either one leaves the deployment unpriced until the config is fixed.
"""
ptu_count: Final = model_info.get("ptu_count")
ptu_count_int: Final = _parsed_ptu_count(model_info)
cost_per_hour: Final = model_info.get("cost_per_ptu_per_hour")
if ptu_count is None or cost_per_hour is None:
if ptu_count_int is None or isinstance(cost_per_hour, bool) or not isinstance(cost_per_hour, (int, float, str)):
return None
try:
ptu_count_int: Final = int(ptu_count)
cost_per_hour_float: Final = float(cost_per_hour)
except (TypeError, ValueError, OverflowError):
return None
if not 0 < ptu_count_int <= ModelInfo.MAX_PTU_COUNT:
except (ValueError, OverflowError):
return None
if not 0 <= cost_per_hour_float <= ModelInfo.MAX_COST_PER_PTU_PER_HOUR:
return None

View file

@ -33,6 +33,7 @@ from litellm._logging import verbose_proxy_logger
from litellm.caching.redis_cache import log_redis_failure
from litellm.constants import DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE, INTERNAL_CALL_ORIGIN_METADATA_KEY
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.azure_ptu_capacity import normalized_tokens
from litellm.litellm_core_utils.prompt_templates.common_utils import (
get_str_from_messages,
)
@ -65,7 +66,7 @@ from litellm.router_utils.add_retry_fallback_headers import (
response_has_hidden_params,
)
from litellm.router_utils.common_utils import resolve_model_group_alias
from litellm.router_utils.ptu_shares import PTUTeamCeiling, team_ptu_ceiling
from litellm.router_utils.ptu_shares import PTUTeamCeiling, model_group_deployments, team_ptu_ceiling
from litellm.types.caching import RedisPipelineIncrementOperation
from litellm.types.llms.openai import BaseLiteLLMOpenAIResponseObject, ResponseAPIUsage
from litellm.types.utils import (
@ -94,6 +95,17 @@ else:
_REQUEST_RATE_LIMIT_DATA: Final = TypeAdapter(Mapping[str, object])
@dataclass(frozen=True, slots=True)
class _ReconciledUsage:
prompt_tokens: int
completion_tokens: int
cached_tokens: int
@property
def billable_input_tokens(self) -> int:
return max(0, self.prompt_tokens - self.cached_tokens)
@dataclass(frozen=True, slots=True)
class RateLimitedModel:
requested: str
@ -119,7 +131,7 @@ def _resolve_ptu_team_ceiling_via_proxy_router(team_id: str, model_group: str) -
if llm_router is None or not is_ptu_cost_attribution_enabled():
return None
return team_ptu_ceiling(llm_router.get_model_list(model_name=model_group) or (), team_id)
return team_ptu_ceiling(model_group_deployments(llm_router.get_model_list() or (), model_group), team_id)
def _sibling_counter_keys(window_key: str) -> tuple[str, str]:
@ -430,9 +442,6 @@ _AUDIO_BYTES_PER_TOKEN: Final = 1600
# on the same project+model simultaneously without colliding on cache keys.
PROJECT_ITPM_DESCRIPTOR_KEY: Final = "model_per_project_itpm"
PROJECT_OTPM_DESCRIPTOR_KEY: Final = "model_per_project_otpm"
# Descriptor "key" for a team's PTU share of a shared Azure provisioned deployment,
# counted in Azure normalized tokens (output weighted by the model's ratio) so it
# never collides with the raw-token "model_per_team" counter on the same team+model.
PTU_TEAM_DESCRIPTOR_KEY: Final = "model_per_team_ptu"
# How long an acquired slot counts toward the in-flight total before it is
# considered leaked (worker crashed without any release callback firing) and
@ -4200,20 +4209,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
return max(0, total_tokens), 0, True
return None
def _resolve_io_token_reconcile_usage(
self,
response_obj: object,
) -> tuple[int, int, bool]:
"""
Resolve ``(billable_input_tokens, completion_tokens, usage_resolved)``
for ITPM/OTPM reconciliation. Cache-read tokens are excluded from
billable input -- Bedrock Mantle doesn't count them toward ITPM --
but they're untouched everywhere else (cost/usage logging still sees
the full prompt token count).
"""
def _resolve_reconciled_usage(self, response_obj: object) -> _ReconciledUsage | None:
"""The prompt, completion, and cache-read token counts a response reports, else None
when it reports no usage at all. Cache-read tokens stay inside ``prompt_tokens`` here;
each consumer decides what they cost it."""
rerank_usage: Final = self._resolve_rerank_token_usage(response_obj)
if rerank_usage is not None:
return rerank_usage
return _ReconciledUsage(prompt_tokens=rerank_usage[0], completion_tokens=rerank_usage[1], cached_tokens=0)
usage: Final = self._response_usage(response_obj)
@ -4226,8 +4228,10 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
else 0
)
if prompt_tokens == 0 and completion_tokens == 0:
return 0, 0, False
return max(0, prompt_tokens - cached_tokens), completion_tokens, True
return None
return _ReconciledUsage(
prompt_tokens=prompt_tokens, completion_tokens=completion_tokens, cached_tokens=cached_tokens
)
if isinstance(usage, ResponseAPIUsage):
response_input_tokens: Final = usage.input_tokens or 0
@ -4236,8 +4240,12 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
usage.input_tokens_details.cached_tokens or 0 if usage.input_tokens_details is not None else 0
)
if response_input_tokens == 0 and response_output_tokens == 0:
return 0, 0, False
return max(0, response_input_tokens - response_cached_tokens), response_output_tokens, True
return None
return _ReconciledUsage(
prompt_tokens=response_input_tokens,
completion_tokens=response_output_tokens,
cached_tokens=response_cached_tokens,
)
if isinstance(usage, Mapping):
raw_prompt_tokens: Final = usage.get("prompt_tokens") or usage.get("input_tokens") or 0
@ -4252,10 +4260,30 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
)
mapped_cached_tokens: Final = raw_cached_tokens if isinstance(raw_cached_tokens, int) else 0
if mapped_prompt_tokens == 0 and mapped_completion_tokens == 0:
return 0, 0, False
return max(0, mapped_prompt_tokens - mapped_cached_tokens), mapped_completion_tokens, True
return None
return _ReconciledUsage(
prompt_tokens=mapped_prompt_tokens,
completion_tokens=mapped_completion_tokens,
cached_tokens=mapped_cached_tokens,
)
return 0, 0, False
return None
def _resolve_io_token_reconcile_usage(
self,
response_obj: object,
) -> tuple[int, int, bool]:
"""
Resolve ``(billable_input_tokens, completion_tokens, usage_resolved)``
for ITPM/OTPM reconciliation. Cache-read tokens are excluded from
billable input -- Bedrock Mantle doesn't count them toward ITPM --
but they're untouched everywhere else (cost/usage logging still sees
the full prompt token count).
"""
usage: Final = self._resolve_reconciled_usage(response_obj)
if usage is None:
return 0, 0, False
return usage.billable_input_tokens, usage.completion_tokens, True
def _build_io_token_reservation_ops(
self,
@ -4617,27 +4645,37 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
reserved_tokens: int,
total_tokens: int,
) -> Sequence[RedisPipelineIncrementOperation]:
"""Settle the team's PTU counter in Azure normalized tokens: uncached input in full plus
output weighted by the model's output-to-input ratio, the way Azure sizes a PTU.
"""Settle the team's PTU counter in Azure normalized tokens: uncached input in full,
cached input at the model's cached ratio, output weighted by its output-to-input
ratio, the way Azure sizes a PTU.
The pre-call reservation was raw estimated tokens, so this is the same reconcile as the
other TPM scopes with a weighted actual; when usage cannot be resolved it charges the raw
total the other scopes charge.
other TPM scopes with a weighted actual; when usage cannot be resolved, or the ceiling
is gone since the reservation was taken, it charges the raw total the other scopes
charge so the reservation is never left standing.
"""
team_id: Final = standard_logging_metadata.get("user_api_key_team_id")
if reconcile_model is None or not isinstance(team_id, str) or not team_id:
return ()
scope: Final = (PTU_TEAM_DESCRIPTOR_KEY, f"{team_id}:{reconcile_model.group}")
ceiling: Final = self._ptu_team_ceiling_resolver(team_id, reconcile_model.group)
if ceiling is None:
if ceiling is None and scope not in reserved_scopes:
return ()
billable_input, completion_tokens, usage_resolved = self._resolve_io_token_reconcile_usage(response_obj)
usage: Final = self._resolve_reconciled_usage(response_obj)
normalized: Final = (
billable_input + round(ceiling.output_to_input_ratio * completion_tokens)
if usage_resolved
round(
normalized_tokens(
ceiling,
prompt_tokens=usage.prompt_tokens,
completion_tokens=usage.completion_tokens,
cache_read_tokens=usage.cached_tokens,
)
)
if ceiling is not None and usage is not None
else total_tokens
)
return self._build_reservation_aware_tpm_ops(
targets=((PTU_TEAM_DESCRIPTOR_KEY, f"{team_id}:{reconcile_model.group}"),),
targets=(scope,),
reserved_scopes=reserved_scopes,
actual_tokens=normalized,
reserved_tokens=reserved_tokens,

View file

@ -201,7 +201,7 @@ from litellm.repositories.verification_token_repository import (
VerificationTokenRepository,
)
from litellm.router import Router
from litellm.router_utils.ptu_shares import model_group_ptu_capacity
from litellm.router_utils.ptu_shares import model_group_deployments, model_group_ptu_capacity
from litellm.types.proxy.auth.auth_checks import UserNotFoundError
from litellm.types.proxy.management_endpoints.common_daily_activity import (
DailySpendMetadata,
@ -6670,7 +6670,9 @@ def _with_ptu_consumption(
return activity
return attach_ptu_hours(
activity,
lambda model_group: model_group_ptu_capacity(llm_router.get_model_list(model_name=model_group) or ()),
lambda model_group: model_group_ptu_capacity(
model_group_deployments(llm_router.get_model_list() or (), model_group)
),
)

View file

@ -454,16 +454,17 @@ async def run_ptu_flat_cost_rollup(
scanned_ids=loaded.scanned_ids,
)
models_processed: Final = len(frozenset(model.model_id for model in ptu_models))
verbose_proxy_logger.info(
"PTU rollup for %s: %d PTU models processed, %d rows written, %d rows failed",
date_str,
len(ptu_models),
models_processed,
rows_written,
rows_failed,
)
return RollupResult(
day=day,
models_processed=len(ptu_models),
models_processed=models_processed,
rows_written=rows_written,
rows_failed=rows_failed,
lapsed=_lapsed_models(ptu_models, run_started),

View file

@ -20,6 +20,7 @@ _DeploymentT = TypeVar("_DeploymentT", bound=Mapping[str, object])
class PTUTeamCeiling:
tpm_limit: int
output_to_input_ratio: float
cached_input_ratio: float
@dataclass(frozen=True, slots=True)
@ -55,8 +56,8 @@ def team_ptu_ceiling(deployments: Sequence[Mapping[str, object]], team_id: str)
up to, else None when the team holds no share on a deployment with a known sizing row.
Two shared deployments of different models in one group are weighted by the larger
output ratio, which over-counts output on the cheaper one rather than under-counting it
on the dearer one.
output and cached-input ratios, which over-counts those tokens on the cheaper one rather
than under-counting them on the dearer one.
"""
priced: Final = tuple(
(shares[team_id], capacity)
@ -70,6 +71,22 @@ def team_ptu_ceiling(deployments: Sequence[Mapping[str, object]], team_id: str)
return PTUTeamCeiling(
tpm_limit=sum(share * capacity.input_tpm_per_ptu 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),
)
def model_group_deployments(deployments: Sequence[_DeploymentT], model_group: str) -> tuple[_DeploymentT, ...]:
"""Every deployment serving ``model_group``: by its own name, or by the public name a
team-scoped deployment keeps in ``model_info.team_public_model_name`` after the router
renames it to a unique internal one."""
return tuple(
deployment
for deployment in deployments
if deployment.get("model_name") == model_group
or (
isinstance(model_info := deployment.get("model_info"), Mapping)
and model_info.get("team_public_model_name") == model_group
)
)

View file

@ -10,7 +10,7 @@ from typing import TYPE_CHECKING, Annotated, Any, ClassVar, Final, Generic, Lite
from zoneinfo import ZoneInfo, ZoneInfoNotFoundError
import httpx
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
from pydantic import BaseModel, ConfigDict, Field, StrictInt, field_validator, model_validator
from typing_extensions import Protocol, ReadOnly, Required, TypedDict, runtime_checkable
from litellm._logging import verbose_logger
@ -260,7 +260,7 @@ class ModelInfo(MirroredPricingParams):
cost_per_ptu_per_hour: float | None = None
ptu_effective_from: datetime.datetime | None = None
ptu_effective_to: datetime.datetime | None = None
ptu_shares: Mapping[str, int] | None = None
ptu_shares: Mapping[str, StrictInt] | None = None
# when tag-based routing's "!" or "&" constraints eliminate every deployment
# in this model group, fall back to the default-tagged pool instead of

View file

@ -55,8 +55,9 @@ def test_a_deployment_prefers_its_declared_base_model_over_its_deployment_name()
def test_a_deployment_falls_back_to_its_litellm_model_when_no_base_model_is_declared():
assert deployment_ptu_capacity({"litellm_params": {"model": "azure/gpt-4o"}}) is AZURE_PTU_CAPACITY["gpt-4o"]
assert deployment_ptu_capacity({"model_info": {"base_model": ""}, "litellm_params": {"model": "azure/gpt-4o"}}) is (
AZURE_PTU_CAPACITY["gpt-4o"]
assert (
deployment_ptu_capacity({"model_info": {"base_model": ""}, "litellm_params": {"model": "azure/gpt-4o"}})
is (AZURE_PTU_CAPACITY["gpt-4o"])
)
@ -72,9 +73,9 @@ def test_output_is_weighted_by_the_models_ratio_and_uncached_input_counts_in_ful
def test_cached_input_is_free_unless_the_row_prices_it():
assert normalized_tokens(_ROW, prompt_tokens=100, completion_tokens=0, cache_read_tokens=60) == pytest.approx(40.0)
assert normalized_tokens(_CACHED_ROW, prompt_tokens=100, completion_tokens=0, cache_read_tokens=60) == pytest.approx(
46.0
)
assert normalized_tokens(
_CACHED_ROW, prompt_tokens=100, completion_tokens=0, cache_read_tokens=60
) == pytest.approx(46.0)
def test_cached_input_never_exceeds_the_prompt_and_negatives_count_as_zero():

View file

@ -395,8 +395,13 @@ def test_a_single_team_reservation_holds_the_whole_count_under_that_team():
{"ptu_shares": {"team-a": 50.5, "team-b": 49.5}},
"ptu_shares must map at least one team_id to a positive whole number of PTUs",
),
({"ptu_shares": {"team-a": True}}, "ptu_shares must map at least one team_id to a positive whole number of PTUs"),
(
{"ptu_shares": {"team-a": True}},
"ptu_shares must map at least one team_id to a positive whole number of PTUs",
),
({"ptu_shares": {"": 100}}, "ptu_shares must map at least one team_id to a positive whole number of PTUs"),
({"ptu_shares": {None: 100}}, "ptu_shares must map at least one team_id to a positive whole number of PTUs"),
({"ptu_shares": {1: 100}}, "ptu_shares must map at least one team_id to a positive whole number of PTUs"),
({"ptu_shares": ["team-a"]}, "ptu_shares must map at least one team_id to a positive whole number of PTUs"),
({"ptu_shares": {"team-a": 60, "team-b": 30}}, "ptu_shares must add up to ptu_count (90 of 100 allocated)"),
({"ptu_shares": {"team-a": 60, "team-b": 50}}, "ptu_shares must add up to ptu_count (110 of 100 allocated)"),
@ -408,6 +413,8 @@ def test_a_single_team_reservation_holds_the_whole_count_under_that_team():
"fractional share",
"boolean share",
"blank team",
"null team",
"numeric team",
"not a mapping",
"shares short of the count",
"shares over the count",
@ -418,6 +425,20 @@ def test_an_incoherent_split_names_its_reason_and_reserves_nothing(override, exp
assert ptu_terms({**_SHARED, **override}) is None
def test_a_whole_count_written_as_a_float_is_checked_against_the_shares_all_the_same():
assert ptu_config_error({**_SHARED, "ptu_count": 100.0, "ptu_shares": {"team-a": 60, "team-b": 30}}) == (
"ptu_shares must add up to ptu_count (90 of 100 allocated)"
)
terms = ptu_terms({**_SHARED, "ptu_count": 100.0})
assert terms is not None
assert terms.ptu_count == 100
def test_a_fractional_count_reserves_nothing():
assert ptu_terms({**_VALID, "ptu_count": 100.5}) is None
assert ptu_terms({**_SHARED, "ptu_count": 100.5}) is None
def test_the_split_is_named_after_the_deployment_when_the_caller_supplies_one():
error = ptu_config_error({**_SHARED, "ptu_shares": {"team-a": 1}}, model_name="gpt-4.1-ptu")

View file

@ -11,6 +11,7 @@ from collections.abc import Iterator, Sequence
from contextlib import contextmanager
from datetime import datetime, timedelta, timezone
from typing import Any, Dict, Final, List, Optional
from unittest.mock import patch
import pytest
from fastapi import HTTPException
@ -20,6 +21,7 @@ from litellm import Router
from litellm.caching.caching import DualCache
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY
from litellm.litellm_core_utils.azure_ptu_capacity import AZURE_PTU_CAPACITY
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
@ -36,6 +38,7 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import (
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
_PROXY_MaxParallelRequestsHandler_v3 as _PROXY_MaxParallelRequestsHandler,
)
from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR
from litellm.proxy.utils import InternalUsageCache, ProxyLogging, hash_token
from litellm.router_utils.ptu_shares import PTUTeamCeiling
from litellm.types.caching import RedisPipelineIncrementOperation
@ -4726,8 +4729,8 @@ async def test_async_data_generator_releases_counter_when_wrapped_v3():
from the outer generator: the counter returns to 0 (not -1), proving the
nested hook does not also refund and there is no double decrement.
"""
from litellm.integrations.custom_logger import CustomLogger
import litellm.proxy.proxy_server as proxy_server
from litellm.integrations.custom_logger import CustomLogger
class _PassthroughIteratorOverride(CustomLogger):
async def async_post_call_streaming_iterator_hook(
@ -7157,14 +7160,14 @@ async def test_success_tpm_accounting_keeps_the_admission_target_after_an_alias_
# --- a team's PTU share on a shared Azure provisioned deployment ---------------------------
def _ptu_ceiling_for(team_id: str, model_group: str, tpm_limit: int, ratio: float):
def _ptu_ceiling_for(team_id: str, model_group: str, tpm_limit: int, ratio: float, cached_ratio: float = 0.0):
calls: list[tuple[str, str]] = []
def resolve(requested_team: str, requested_group: str) -> PTUTeamCeiling | None:
calls.append((requested_team, requested_group))
if (requested_team, requested_group) != (team_id, model_group):
return None
return PTUTeamCeiling(tpm_limit=tpm_limit, output_to_input_ratio=ratio)
return PTUTeamCeiling(tpm_limit=tpm_limit, output_to_input_ratio=ratio, cached_input_ratio=cached_ratio)
return resolve, calls
@ -7177,12 +7180,16 @@ def _ptu_request(model: str = "test-model") -> dict:
async def test_a_teams_ptu_share_is_a_hard_tpm_ceiling_on_the_shared_model():
cache = DualCache()
resolve, _ = _ptu_ceiling_for("t", "test-model", tpm_limit=500, ratio=4.0)
handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache), ptu_team_ceiling_resolver=resolve)
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(cache), ptu_team_ceiling_resolver=resolve
)
key = UserAPIKeyAuth(api_key=hash_token("sk-ptu"), team_id="t")
await handler.async_pre_call_hook(user_api_key_dict=key, cache=cache, data=_ptu_request(), call_type="acompletion")
with pytest.raises(HTTPException) as exc:
await handler.async_pre_call_hook(user_api_key_dict=key, cache=cache, data=_ptu_request(), call_type="acompletion")
await handler.async_pre_call_hook(
user_api_key_dict=key, cache=cache, data=_ptu_request(), call_type="acompletion"
)
assert exc.value.status_code == 429
assert "model_per_team_ptu" in str(exc.value.detail)
@ -7193,7 +7200,9 @@ async def test_a_teams_ptu_share_is_a_hard_tpm_ceiling_on_the_shared_model():
async def test_a_ptu_ceiling_on_one_model_leaves_the_teams_other_models_alone():
cache = DualCache()
resolve, calls = _ptu_ceiling_for("t", "test-model", tpm_limit=500, ratio=4.0)
handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache), ptu_team_ceiling_resolver=resolve)
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(cache), ptu_team_ceiling_resolver=resolve
)
key = UserAPIKeyAuth(api_key=hash_token("sk-ptu"), team_id="t")
for _ in range(3):
@ -7208,28 +7217,102 @@ async def test_a_ptu_ceiling_on_one_model_leaves_the_teams_other_models_alone():
async def test_a_team_without_a_share_and_a_key_without_a_team_get_no_ptu_ceiling():
cache = DualCache()
resolve, calls = _ptu_ceiling_for("t", "test-model", tpm_limit=500, ratio=4.0)
handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache), ptu_team_ceiling_resolver=resolve)
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(cache), ptu_team_ceiling_resolver=resolve
)
other_team = UserAPIKeyAuth(api_key=hash_token("sk-other"), team_id="u")
no_team = UserAPIKeyAuth(api_key=hash_token("sk-no-team"))
for _ in range(3):
await handler.async_pre_call_hook(user_api_key_dict=other_team, cache=cache, data=_ptu_request(), call_type="acompletion")
await handler.async_pre_call_hook(user_api_key_dict=no_team, cache=cache, data=_ptu_request(), call_type="acompletion")
for caller in (other_team, no_team):
await handler.async_pre_call_hook(
user_api_key_dict=caller, cache=cache, data=_ptu_request(), call_type="acompletion"
)
assert set(calls) == {("u", "test-model")}
def _shared_ptu_router(model_group: str) -> Router:
return Router(
model_list=[
{
"model_name": model_group,
"litellm_params": {"model": "azure/gpt-4.1", "api_key": "sk-ptu", "api_base": "https://ptu.example"},
"model_info": {
"id": "shared-ptu",
"base_model": "azure/gpt-4.1",
"ptu_count": 1,
"cost_per_ptu_per_hour": 1.0,
"ptu_effective_from": "2026-01-01T00:00:00Z",
"ptu_shares": {"t": 1},
},
}
]
)
def _two_thirds_of_a_ptu_minute() -> dict:
return {**_ptu_request(), "max_tokens": AZURE_PTU_CAPACITY["gpt-4.1"].input_tpm_per_ptu * 2 // 3}
@pytest.mark.asyncio
async def test_the_proxy_router_turns_a_teams_share_into_its_ceiling_when_attribution_is_on(monkeypatch):
"""With no resolver injected the ceiling comes from the proxy router's own deployments: one
PTU of gpt-4.1 a minute, so two requests each reserving two thirds of it are one too many."""
monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true")
cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache))
key = UserAPIKeyAuth(api_key=hash_token("sk-ptu"), team_id="t")
with patch("litellm.proxy.proxy_server.llm_router", _shared_ptu_router("test-model")):
await handler.async_pre_call_hook(
user_api_key_dict=key, cache=cache, data=_two_thirds_of_a_ptu_minute(), call_type="acompletion"
)
with pytest.raises(HTTPException) as exc:
await handler.async_pre_call_hook(
user_api_key_dict=key, cache=cache, data=_two_thirds_of_a_ptu_minute(), call_type="acompletion"
)
assert exc.value.status_code == 429
assert "model_per_team_ptu" in str(exc.value.detail)
@pytest.mark.asyncio
async def test_the_proxy_router_sets_no_ceiling_while_attribution_is_off(monkeypatch):
monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False)
cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache))
key = UserAPIKeyAuth(api_key=hash_token("sk-ptu"), team_id="t")
with patch("litellm.proxy.proxy_server.llm_router", _shared_ptu_router("test-model")):
for _ in range(3):
await handler.async_pre_call_hook(
user_api_key_dict=key, cache=cache, data=_two_thirds_of_a_ptu_minute(), call_type="acompletion"
)
assert not any("model_per_team_ptu" in cache_key for cache_key in cache.in_memory_cache.cache_dict)
def _ptu_success_kwargs() -> dict:
return {
"standard_logging_object": {"metadata": {"user_api_key_hash": hash_token("sk-ptu"), "user_api_key_team_id": "t"}},
"litellm_params": {"metadata": {"model_group": "test-model", "user_api_key_metadata": {}, "user_api_key_team_metadata": {}}},
"standard_logging_object": {
"metadata": {"user_api_key_hash": hash_token("sk-ptu"), "user_api_key_team_id": "t"}
},
"litellm_params": {
"metadata": {"model_group": "test-model", "user_api_key_metadata": {}, "user_api_key_team_metadata": {}}
},
"model": "test-model",
}
def _ptu_response(usage: Usage) -> ModelResponse:
return ModelResponse(
id="ptu-share", object="chat.completion", created=int(datetime.now().timestamp()), model="test-model", usage=usage, choices=[]
id="ptu-share",
object="chat.completion",
created=int(datetime.now().timestamp()),
model="test-model",
usage=usage,
choices=[],
)
@ -7275,6 +7358,59 @@ def test_cached_input_is_not_charged_to_the_ptu_counter():
assert _ptu_increment(handler, ops) == 60 + 4 * 50
def test_cached_input_is_charged_at_the_models_cached_ratio():
"""40 of the 100 input tokens were cache reads; at a tenth each they are 4 normalized tokens
beside the 60 uncached ones and the 200 for 50 outputs at 4:1."""
resolve, _ = _ptu_ceiling_for("t", "test-model", tpm_limit=500, ratio=4.0, cached_ratio=0.1)
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache()), ptu_team_ceiling_resolver=resolve
)
response = _ptu_response(
Usage(
prompt_tokens=100,
completion_tokens=50,
total_tokens=150,
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=40),
)
)
ops = handler._build_success_event_pipeline_operations(
kwargs=_ptu_success_kwargs(), response_obj=response, rate_limit_type="output"
)
assert _ptu_increment(handler, ops) == 60 + 4 + 4 * 50
@pytest.mark.asyncio
async def test_a_reservation_is_settled_even_after_the_teams_share_is_gone():
"""The share can be removed between admission and completion; the reserved tokens still
come off the counter instead of standing in the window."""
ceiling: dict[str, PTUTeamCeiling | None] = {
"current": PTUTeamCeiling(tpm_limit=500, output_to_input_ratio=4.0, cached_input_ratio=0.0)
}
cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(cache),
ptu_team_ceiling_resolver=lambda _team, _group: ceiling["current"],
)
key = UserAPIKeyAuth(api_key=hash_token("sk-ptu"), team_id="t")
await handler.async_pre_call_hook(user_api_key_dict=key, cache=cache, data=_ptu_request(), call_type="acompletion")
stash = get_request_stash()
assert stash is not None
assert ("model_per_team_ptu", "t:test-model") in stash.reserved_scopes
assert stash.reserved_tokens > 150
ceiling["current"] = None
ops = handler._build_success_event_pipeline_operations(
kwargs=_ptu_success_kwargs(),
response_obj=_ptu_response(Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150)),
rate_limit_type="total",
)
assert _ptu_increment(handler, ops) == 150 - stash.reserved_tokens
def test_usage_that_only_reports_a_total_charges_that_total_to_the_ptu_counter():
resolve, _ = _ptu_ceiling_for("t", "test-model", tpm_limit=500, ratio=4.0)
handler = _PROXY_MaxParallelRequestsHandler(

View file

@ -9,31 +9,30 @@ from unittest.mock import patch as patch_ctx
import pytest
from fastapi import HTTPException
from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token
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,
LitellmUserRoles,
ReconcileOutcome,
UserAPIKeyAuth,
)
from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token
from litellm.litellm_core_utils.ptu_pricing import ptu_terms
from litellm.proxy.auth.auth_checks import _is_model_cost_zero
from litellm.llms.gemini.cost_calculator import cost_per_web_search_request
from litellm.proxy.management_endpoints.model_management_endpoints import (
_PTU_ZEROED_PRICING_FIELDS,
_SEARCH_CONTEXT_SIZES,
_is_nonzero_price,
_merged_ptu_model_info,
_update_team_model_in_db,
_ptu_priced_deployment,
_ptu_zeroed_pricing,
_raise_if_ptu_cost_attribution_disabled,
_update_team_model_in_db,
_validate_ptu_model_info,
add_new_model,
update_db_model,
)
from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR
from litellm.types.utils import PromptTokensDetailsWrapper
from litellm.router import Router
from litellm.types.router import (
SPECIAL_MODEL_INFO_PARAMS,
@ -43,7 +42,7 @@ from litellm.types.router import (
updateDeployment,
updateLiteLLMParams,
)
from litellm.types.utils import Usage
from litellm.types.utils import PromptTokensDetailsWrapper, Usage
async def _passthrough_row(update_data):
@ -116,6 +115,18 @@ def test_model_info_allows_partial_delta_for_patch():
assert info.cost_per_ptu_per_hour is None
@pytest.mark.parametrize("share", [True, "2", 2.0])
def test_model_info_rejects_a_share_that_is_not_a_whole_number(share):
with pytest.raises(ValueError, match="ptu_shares"):
ModelInfo(id="x", ptu_shares={"team-a": share})
def test_model_info_keeps_whole_number_shares_and_refuses_a_fractional_count():
assert ModelInfo(id="x", ptu_shares={"team-a": 2}).ptu_shares == {"team-a": 2}
with pytest.raises(ValueError, match="ptu_count"):
ModelInfo(id="x", ptu_count=100.5)
def test_validate_helper_no_ptu_is_noop():
_validate_ptu_model_info({"team_id": "t"})

View file

@ -1,9 +1,9 @@
import asyncio
import json
from collections.abc import Sequence
from contextlib import asynccontextmanager, contextmanager
from datetime import datetime, timezone
from types import SimpleNamespace
from collections.abc import Sequence
from typing import Final, Optional, cast
from unittest.mock import AsyncMock, MagicMock, PropertyMock, call, patch
@ -15,6 +15,7 @@ from pydantic import ValidationError
from litellm._uuid import uuid
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.azure_ptu_capacity import AZURE_PTU_CAPACITY
from litellm.proxy._types import (
LiteLLM_BudgetTable,
LiteLLM_BudgetTableFull,
@ -17106,6 +17107,9 @@ def test_list_team_v2_answers_503_no_db_connection_when_the_callers_user_read_hi
# --- PTU-equivalent consumption on /team/daily/activity ------------------------------------
_ONE_PTU_HOUR_OF_INPUT: Final = AZURE_PTU_CAPACITY["gpt-4.1"].normalized_tokens_per_ptu_hour
def _ptu_activity_page():
from litellm.types.proxy.management_endpoints.common_daily_activity import (
BreakdownMetrics,
@ -17116,8 +17120,9 @@ def _ptu_activity_page():
SpendMetrics,
)
tokens: Final = _ONE_PTU_HOUR_OF_INPUT
metrics = SpendMetrics(
prompt_tokens=180_000, completion_tokens=0, total_tokens=180_000, api_requests=3, successful_requests=3
prompt_tokens=tokens, completion_tokens=0, total_tokens=tokens, api_requests=3, successful_requests=3
)
return SpendAnalyticsPaginatedResponse(
results=[
@ -17129,7 +17134,7 @@ def _ptu_activity_page():
),
)
],
metadata=DailySpendMetadata(total_tokens=180_000, total_api_requests=3, total_successful_requests=3),
metadata=DailySpendMetadata(total_tokens=tokens, total_api_requests=3, total_successful_requests=3),
)
@ -17157,9 +17162,9 @@ def _shared_ptu_router() -> Router:
async def test_team_daily_activity_reports_ptu_hours_only_while_attribution_is_on(
mock_db_client, mock_admin_auth, monkeypatch, attribution_enabled
):
"""One PTU serves 3,000 input tokens per minute on gpt-4.1, so 180,000 uncached
input tokens are one PTU-hour; the figure appears beside tokens only once the
PTU flag is on, and the token totals are untouched either way."""
"""An hour of one PTU's input rate on gpt-4.1, in uncached input tokens, is one PTU-hour;
the figure appears beside tokens only once the PTU flag is on, and the token totals are
untouched either way."""
from litellm.proxy.management_endpoints.team_endpoints import get_team_daily_activity
monkeypatch.setenv("LITELLM_ENABLE_PTU_COST_ATTRIBUTION", "True" if attribution_enabled else "False")
@ -17186,8 +17191,58 @@ async def test_team_daily_activity_reports_ptu_hours_only_while_attribution_is_o
assert result.metadata.total_ptu_hours == expected_ptu_hours
assert result.results[0].metrics.ptu_hours == expected_ptu_hours
assert result.results[0].breakdown.model_groups["gpt-4.1-ptu"].metrics.ptu_hours == expected_ptu_hours
assert result.metadata.total_tokens == 180_000
assert result.results[0].metrics.total_tokens == 180_000
assert result.metadata.total_tokens == _ONE_PTU_HOUR_OF_INPUT
assert result.results[0].metrics.total_tokens == _ONE_PTU_HOUR_OF_INPUT
@pytest.mark.asyncio
async def test_team_daily_activity_sizes_a_team_scoped_deployment_by_its_public_name(
mock_db_client, mock_admin_auth, monkeypatch
):
"""A deployment registered for one team is renamed to a unique internal name and keeps the
name callers use in ``team_public_model_name``, which is the name the activity rows carry,
so its sizing row is still found."""
from litellm.proxy.management_endpoints.team_endpoints import get_team_daily_activity
monkeypatch.setenv("LITELLM_ENABLE_PTU_COST_ATTRIBUTION", "True")
mock_db_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[])
page = _ptu_activity_page()
team_scoped_router: Final = Router(
model_list=[
{
"model_name": "gpt-4.1-ptu-3f9c1b",
"litellm_params": {"model": "azure/gpt-4.1", "api_key": "sk-ptu", "api_base": "https://ptu.example"},
"model_info": {
"id": "team-a-ptu",
"base_model": "azure/gpt-4.1",
"team_id": "team-a",
"team_public_model_name": "gpt-4.1-ptu",
"ptu_count": 50,
"cost_per_ptu_per_hour": 1.0,
"ptu_effective_from": "2026-01-01T00:00:00Z",
},
}
]
)
with (
patch("litellm.proxy.management_endpoints.team_endpoints.get_daily_activity", AsyncMock(return_value=page)),
patch("litellm.proxy.proxy_server.llm_router", team_scoped_router),
):
result = await get_team_daily_activity(
team_ids="team-a",
start_date="2026-09-23",
end_date="2026-09-24",
model=None,
api_key=None,
page=1,
page_size=10,
exclude_team_ids=None,
user_api_key_dict=mock_admin_auth,
)
assert result.metadata.total_ptu_hours == 1.0
assert result.results[0].breakdown.model_groups["gpt-4.1-ptu"].metrics.ptu_hours == 1.0
def test_team_export_csv_columns_match_the_dashboard_client_layout():

View file

@ -31,3 +31,15 @@ def test_reads_the_env_var_on_every_call(monkeypatch):
monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true")
assert is_ptu_cost_attribution_enabled() is True
def test_reads_the_process_environment_and_never_a_secret_manager(monkeypatch):
"""The flag is checked on every request and every rollup, so it never costs a round trip to
a hosted secret manager even when one is configured for reads."""
def refuse(*args: object, **kwargs: object) -> object:
raise AssertionError("secret manager consulted")
monkeypatch.setattr("litellm.secret_managers.main.get_secret", refuse)
monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true")
assert is_ptu_cost_attribution_enabled() is True

View file

@ -2119,7 +2119,10 @@ async def test_rollup_splits_a_shared_deployments_flat_cost_by_share():
result = await run_ptu_flat_cost_rollup(prisma, target_date=DAY)
assert result.rows_written == 2
created = {call.kwargs["data"]["create"]["team_id"]: call.kwargs["data"]["create"] for call in table.upsert.await_args_list}
assert result.models_processed == 1
created = {
call.kwargs["data"]["create"]["team_id"]: call.kwargs["data"]["create"] for call in table.upsert.await_args_list
}
assert created["team-a"]["ptu_flat_cost"] == pytest.approx(720.0)
assert created["team-b"]["ptu_flat_cost"] == pytest.approx(480.0)
assert sum(row["ptu_flat_cost"] for row in created.values()) == pytest.approx(50 * 1.0 * 24)

View file

@ -6,6 +6,7 @@ from litellm.litellm_core_utils.azure_ptu_capacity import AZURE_PTU_CAPACITY
from litellm.router_utils.ptu_shares import (
PTUTeamCeiling,
filter_ptu_shared_deployments,
model_group_deployments,
model_group_ptu_capacity,
ptu_capacity_warning,
team_ptu_ceiling,
@ -13,6 +14,7 @@ from litellm.router_utils.ptu_shares import (
_GPT41: Final = AZURE_PTU_CAPACITY["gpt-4.1"]
_GPT4O: Final = AZURE_PTU_CAPACITY["gpt-4o"]
_GPT6SOL: Final = AZURE_PTU_CAPACITY["gpt-6-sol"]
_SHARES: Final = {"team-a": 30, "team-b": 20}
@ -77,7 +79,9 @@ def test_a_single_team_deployment_and_a_malformed_share_map_are_not_filtered_her
def test_a_share_converts_to_the_models_input_tpm_per_ptu():
ceiling: Final = team_ptu_ceiling([_shared()], "team-a")
assert ceiling == PTUTeamCeiling(
tpm_limit=30 * _GPT41.input_tpm_per_ptu, output_to_input_ratio=_GPT41.output_to_input_ratio
tpm_limit=30 * _GPT41.input_tpm_per_ptu,
output_to_input_ratio=_GPT41.output_to_input_ratio,
cached_input_ratio=_GPT41.cached_input_ratio,
)
@ -89,6 +93,32 @@ def test_shares_across_deployments_add_up_and_the_larger_output_ratio_wins():
assert ceiling.output_to_input_ratio == max(_GPT41.output_to_input_ratio, _GPT4O.output_to_input_ratio)
def test_the_larger_cached_input_ratio_wins_across_deployments():
"""A team sharing two models is weighted by the one that charges more for cache reads,
whichever order the deployments come in."""
gpt6sol: Final = _shared(model="azure/gpt-6-sol", shares={"team-a": 10}, deployment_id="shared-6")
ceiling: Final = team_ptu_ceiling([_shared(), gpt6sol], "team-a")
assert ceiling is not None
assert _GPT41.cached_input_ratio < _GPT6SOL.cached_input_ratio
assert ceiling.cached_input_ratio == _GPT6SOL.cached_input_ratio
def test_a_group_is_served_by_name_or_by_a_team_scoped_deployments_public_name():
"""A deployment registered for one team is renamed to a unique internal name and keeps
the name callers use in ``team_public_model_name``."""
team_scoped: Final = {
"model_name": "gpt-4.1-ptu-3f9c1b",
"litellm_params": {"model": "azure/gpt-4.1"},
"model_info": {"id": "team-scoped", "team_id": "team-a", "team_public_model_name": "gpt-4.1-ptu"},
}
other: Final = {"model_name": "other", "litellm_params": {"model": "azure/gpt-4o"}, "model_info": {"id": "other"}}
deployments: Final = [team_scoped, _OPEN, other]
served: Final = model_group_deployments(deployments, "gpt-4.1-ptu")
assert [d["model_info"]["id"] for d in served] == ["team-scoped", "open"]
assert model_group_deployments(deployments, "gpt-4.1-ptu-3f9c1b") == (team_scoped,)
assert model_group_deployments(deployments, "missing") == ()
def test_no_share_or_no_sizing_row_sets_no_ceiling():
assert team_ptu_ceiling([_shared()], "team-c") is None
assert team_ptu_ceiling([_shared(model="azure/unknown-deployment")], "team-a") is None
@ -112,4 +142,5 @@ def test_a_reserved_deployment_without_a_sizing_row_is_warned_about_by_name():
def test_a_sized_reservation_and_an_unreserved_deployment_raise_no_warning():
assert ptu_capacity_warning("gpt-4.1-ptu", _shared()) is None
assert ptu_capacity_warning("gpt-4.1-ptu", _single_team()) is None
assert ptu_capacity_warning("gpt-4.1-ptu", {**_OPEN, "litellm_params": {"model": "azure/my-ptu-deployment"}}) is None
unsized_open: Final = {**_OPEN, "litellm_params": {"model": "azure/my-ptu-deployment"}}
assert ptu_capacity_warning("gpt-4.1-ptu", unsized_open) is None