fix(proxy): reserve PTU ceilings in normalized tokens, keep share routing flag-independent, and warn only shared or Azure reservations

This commit is contained in:
mateo-berri 2026-09-24 15:39:14 -07:00
parent 304e9d2ca0
commit 0d80930df8
12 changed files with 289 additions and 53 deletions

View file

@ -33,12 +33,12 @@ 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,
)
from litellm.litellm_core_utils.ptu_pricing import is_ptu_cost_attribution_enabled
from litellm.litellm_core_utils.token_counter import offload_token_count
from litellm.llms.azure.ptu_capacity import normalized_tokens
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_utils import (
ESTIMATED_OUTPUT_TOKENS_FIELD,
@ -592,6 +592,8 @@ class RequestRateLimiterStash:
reserved_tokens: int = 0
reserved_model: RateLimitedModel | None = None
reserved_scopes: frozenset[tuple[str, str]] = field(default_factory=frozenset)
ptu_ceiling: PTUTeamCeiling | None = None
ptu_reserved_tokens: int = 0
itpm_reserved_tokens: int = 0
itpm_reserved_scopes: frozenset[tuple[str, str]] = field(default_factory=frozenset)
itpm_reserved_window_identities: frozenset[tuple[str, str, Literal["redis", "local"]]] = field(
@ -973,6 +975,30 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
return total_estimated
def _estimate_ptu_tokens_for_request(
self,
ceiling: PTUTeamCeiling | None,
data: dict,
min_configured_tpm_limit: int | None,
call_type: str | None,
configured_output_tokens: int | None,
raw_estimate: int,
) -> int:
"""The team PTU ceiling counts Azure normalized tokens, so its reservation weighs the
output budget the way the ceiling does instead of the raw sum the other scopes reserve."""
if ceiling is None:
return raw_estimate
estimated_input_tokens, max_tokens_estimate = self._estimate_input_and_output_tokens(
data=data,
min_configured_tpm_limit=min_configured_tpm_limit,
call_type=call_type,
configured_output_tokens=configured_output_tokens,
)
normalized: Final = normalized_tokens(
ceiling, prompt_tokens=estimated_input_tokens, completion_tokens=max_tokens_estimate
)
return max(round(normalized), 1)
def _estimate_input_and_output_tokens(
self,
data: object,
@ -2184,11 +2210,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
descriptors: list[RateLimitDescriptor],
estimated_tokens: int,
parent_otel_span: Span | None = None,
scope_estimates: Mapping[str, int] = MappingProxyType({}),
) -> RateLimitResponse:
"""
Reserve ``estimated_tokens`` against every TPM-bearing descriptor
BEFORE the upstream call, so concurrent requests cannot all observe
"under limit" before any of them increments the counter.
``scope_estimates`` replaces that amount per descriptor key for a
scope counted in other units, such as the team PTU ceiling.
Thin wrapper around ``atomic_check_and_increment_by_n``: builds a
TPM-only descriptor/increment list and delegates the all-or-nothing
@ -2215,7 +2244,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
return RateLimitResponse(overall_code="OK", statuses=[])
increments: Final[list[dict[Literal["requests", "tokens"], int]]] = [
{"tokens": estimated_tokens} for _ in tpm_descriptors
{"tokens": scope_estimates.get(d["key"], estimated_tokens)} for d in tpm_descriptors
]
return await self.atomic_check_and_increment_by_n(
descriptors=tpm_descriptors,
@ -3061,6 +3090,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
ceiling: Final = self._ptu_team_ceiling_resolver(user_api_key_dict.team_id, model.group)
if ceiling is None:
return
stash: Final = get_request_stash()
if stash is not None:
stash.ptu_ceiling = ceiling
descriptors.append(
RateLimitDescriptor(
key=PTU_TEAM_DESCRIPTOR_KEY,
@ -3482,10 +3514,15 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
# (still-stashed) reservation and refunds it again.
if tpm_reservation_amount > 0:
await self._refund_reserved_tokens(
scopes=tpm_reservation_scopes,
scopes=tuple(scope for scope in tpm_reservation_scopes if scope[0] != PTU_TEAM_DESCRIPTOR_KEY),
amount=tpm_reservation_amount,
parent_otel_span=user_api_key_dict.parent_otel_span,
)
await self._refund_reserved_tokens(
scopes=tuple(scope for scope in tpm_reservation_scopes if scope[0] == PTU_TEAM_DESCRIPTOR_KEY),
amount=stash.ptu_reserved_tokens,
parent_otel_span=user_api_key_dict.parent_otel_span,
)
stash.reservation_released = True
await self._release_stashed_parallel_slot(stash, user_api_key_dict.parent_otel_span)
self._handle_rate_limit_error(
@ -3795,10 +3832,19 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
min_configured_tpm_limit,
)
ptu_estimated_tokens: Final = self._estimate_ptu_tokens_for_request(
ceiling=stash.ptu_ceiling,
data=data,
min_configured_tpm_limit=min_configured_tpm_limit,
call_type=call_type,
configured_output_tokens=configured_output_tokens,
raw_estimate=estimated_tokens,
)
tpm_response: Final = await self.reserve_tpm_tokens(
descriptors=descriptors,
estimated_tokens=estimated_tokens,
parent_otel_span=user_api_key_dict.parent_otel_span,
scope_estimates=MappingProxyType({PTU_TEAM_DESCRIPTOR_KEY: ptu_estimated_tokens}),
)
if tpm_response["overall_code"] == "OVER_LIMIT":
@ -3814,6 +3860,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
# the (actual - reserved) delta to those — unreserved
# scopes get charged the full actual usage instead.
stash.reserved_tokens = estimated_tokens
stash.ptu_reserved_tokens = ptu_estimated_tokens
stash.reserved_model = self._rate_limited_model(requested_model)
stash.reserved_scopes = frozenset(
(d["key"], d["value"])
@ -4629,7 +4676,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
response_obj=response_obj,
reconcile_model=reconcile_model,
reserved_scopes=reserved_scopes,
reserved_tokens=reserved_tokens,
reserved_ceiling=stash.ptu_ceiling if stash is not None else None,
reserved_tokens=stash.ptu_reserved_tokens if stash is not None else 0,
total_tokens=total_tokens,
)
)
@ -4642,6 +4690,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
response_obj: object,
reconcile_model: RateLimitedModel | None,
reserved_scopes: Set[tuple[str, str]],
reserved_ceiling: PTUTeamCeiling | None,
reserved_tokens: int,
total_tokens: int,
) -> Sequence[RedisPipelineIncrementOperation]:
@ -4649,38 +4698,44 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
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, 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.
The reservation was taken in those units against the ceiling admission resolved, so that
ceiling settles it even after the share changed or went away; when usage cannot be
resolved 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)
ceiling: Final = (
reserved_ceiling
if reserved_ceiling is not None
else self._ptu_team_ceiling_resolver(team_id, reconcile_model.group)
)
if ceiling is None and scope not in reserved_scopes:
return ()
usage: Final = self._resolve_reconciled_usage(response_obj)
normalized: Final = (
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=(scope,),
reserved_scopes=reserved_scopes,
actual_tokens=normalized,
actual_tokens=self._ptu_settlement_tokens(
ceiling, self._resolve_reconciled_usage(response_obj), total_tokens
),
reserved_tokens=reserved_tokens,
)
@staticmethod
def _ptu_settlement_tokens(ceiling: PTUTeamCeiling | None, usage: _ReconciledUsage | None, raw_total: int) -> int:
if ceiling is None or usage is None:
return raw_total
return round(
normalized_tokens(
ceiling,
prompt_tokens=usage.prompt_tokens,
completion_tokens=usage.completion_tokens,
cache_read_tokens=usage.cached_tokens,
)
)
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
"""
Update TPM usage on successful API calls by incrementing counters using pipeline
@ -4786,9 +4841,15 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
statuses=statuses,
)
def _recovered_partial_usage_tokens(self, source: Mapping[str, object]) -> tuple[int, int, int]:
def _recovered_partial_usage(self, source: Mapping[str, object]) -> Usage | None:
usage: Final = source.get("combined_usage_object")
if not isinstance(usage, Usage) or (usage.completion_tokens or 0) <= 0:
return None
return usage
def _recovered_partial_usage_tokens(self, source: Mapping[str, object]) -> tuple[int, int, int]:
usage: Final = self._recovered_partial_usage(source)
if usage is None:
return 0, 0, 0
billable_input, completion_tokens, _ = self._resolve_io_token_reconcile_usage(usage)
return (
@ -4797,6 +4858,26 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
completion_tokens,
)
def _build_ptu_failure_settlement_ops(
self, stash: RequestRateLimiterStash, source: Mapping[str, object], raw_actual_tokens: int
) -> Sequence[RedisPipelineIncrementOperation]:
"""Settle the team PTU reservation on failure in the normalized tokens it was taken in:
at the recovered partial usage when there is one, else a full refund."""
ptu_scopes: Final = tuple(scope for scope in stash.reserved_scopes if scope[0] == PTU_TEAM_DESCRIPTOR_KEY)
if not ptu_scopes:
return ()
usage: Final = self._recovered_partial_usage(source)
return self._build_reservation_aware_tpm_ops(
targets=ptu_scopes,
reserved_scopes=stash.reserved_scopes,
actual_tokens=self._ptu_settlement_tokens(
stash.ptu_ceiling,
self._resolve_reconciled_usage(usage) if usage is not None else None,
raw_actual_tokens,
),
reserved_tokens=stash.ptu_reserved_tokens,
)
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
"""
On failure: decrement max_parallel_requests and refund the upfront
@ -4838,12 +4919,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
# refund there would drive their counter negative.
pipeline_operations.extend(
self._build_reservation_aware_tpm_ops(
targets=list(stash.reserved_scopes),
targets=tuple(scope for scope in stash.reserved_scopes if scope[0] != PTU_TEAM_DESCRIPTOR_KEY),
reserved_scopes=stash.reserved_scopes,
actual_tokens=tpm_actual,
reserved_tokens=reserved_tokens,
)
)
pipeline_operations.extend(self._build_ptu_failure_settlement_ops(stash, kwargs, tpm_actual))
# Settle project ITPM/OTPM reservations the same way: at the
# recovered partial usage, or a full refund when there is none.
@ -5034,11 +5116,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
tpm_actual, itpm_actual, otpm_actual = self._recovered_partial_usage_tokens(request_data)
combined_ops: Final = (
self._build_reservation_aware_tpm_ops(
targets=tuple(stash.reserved_scopes),
reserved_scopes=stash.reserved_scopes,
actual_tokens=tpm_actual,
reserved_tokens=reserved_tokens,
(
*self._build_reservation_aware_tpm_ops(
targets=tuple(scope for scope in stash.reserved_scopes if scope[0] != PTU_TEAM_DESCRIPTOR_KEY),
reserved_scopes=stash.reserved_scopes,
actual_tokens=tpm_actual,
reserved_tokens=reserved_tokens,
),
*self._build_ptu_failure_settlement_ops(stash, request_data, tpm_actual),
)
if reserved_tokens > 0
else ()

View file

@ -9,7 +9,7 @@ from collections.abc import Callable
from types import MappingProxyType
from typing import Final
from litellm.litellm_core_utils.azure_ptu_capacity import PTUCapacity, normalized_tokens, ptu_hours
from litellm.llms.azure.ptu_capacity import PTUCapacity, normalized_tokens, ptu_hours
from litellm.types.proxy.management_endpoints.common_daily_activity import (
DailySpendData,
MetricWithMetadata,

View file

@ -12793,8 +12793,6 @@ class Router:
model=model,
llm_provider="",
)
if not is_ptu_cost_attribution_enabled():
return result.deployments
shared: Final = filter_ptu_shared_deployments(result.deployments, request_team_id)
if shared.withheld and len(shared.deployments) == 0:
raise litellm.BadRequestError(

View file

@ -10,8 +10,11 @@ from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from typing import Final, Generic, TypeVar
from litellm.litellm_core_utils.azure_ptu_capacity import PTUCapacity, deployment_ptu_capacity
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})
_DeploymentT = TypeVar("_DeploymentT", bound=Mapping[str, object])
@ -105,14 +108,38 @@ 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."""
"""Why this reserved deployment's tokens cannot be converted to PTUs, else None.
Only a deployment that declares shares (whose ceilings need a sizing row) or one served by
Azure (whose PTU hours need one) is worth warning about; a single-team reservation on another
provider only ever used the flat-cost rollup, which needs no sizing.
"""
model_info: Final = deployment.get("model_info")
if not isinstance(model_info, Mapping) or ptu_terms(model_info) is None:
return None
if deployment_ptu_capacity(deployment) is not None:
return None
return (
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"
)
has_shares: Final = model_info.get("ptu_shares") is not None
if has_shares:
return (
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):
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"
)
return None

View file

@ -4,7 +4,7 @@ from typing import Final
import pytest
from litellm.litellm_core_utils.azure_ptu_capacity import (
from litellm.llms.azure.ptu_capacity import (
AZURE_PTU_CAPACITY,
PTUCapacity,
azure_ptu_capacity,

View file

@ -21,7 +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.llms.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 (
@ -4500,7 +4500,7 @@ async def test_tpm_over_limit_rejection_releases_parallel_slot_v3(monkeypatch):
)
counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests"
async def over_limit_reservation(descriptors, estimated_tokens, parent_otel_span=None):
async def over_limit_reservation(descriptors, estimated_tokens, parent_otel_span=None, **_kwargs):
return {
"overall_code": "OVER_LIMIT",
"statuses": [
@ -7179,7 +7179,7 @@ def _ptu_request(model: str = "test-model") -> dict:
@pytest.mark.asyncio
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)
resolve, _ = _ptu_ceiling_for("t", "test-model", tpm_limit=2000, ratio=4.0)
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(cache), ptu_team_ceiling_resolver=resolve
)
@ -7252,7 +7252,55 @@ def _shared_ptu_router(model_group: str) -> Router:
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}
"""An output budget worth two thirds of a gpt-4.1 PTU minute once weighted at the model's
output-to-input ratio, the way the ceiling counts it."""
capacity = AZURE_PTU_CAPACITY["gpt-4.1"]
return {**_ptu_request(), "max_tokens": int(capacity.input_tpm_per_ptu * 2 / 3 / capacity.output_to_input_ratio)}
@pytest.mark.asyncio
async def test_the_reservation_weighs_output_the_way_the_ceiling_does():
"""300 output tokens are 1200 normalized tokens at 4:1, over a 1000-token ceiling the raw
301-token estimate would clear; the same request at 1:1 is admitted."""
key = UserAPIKeyAuth(api_key=hash_token("sk-ptu"), team_id="t")
unweighted_cache = DualCache()
unweighted = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(unweighted_cache),
ptu_team_ceiling_resolver=_ptu_ceiling_for("t", "test-model", tpm_limit=1000, ratio=1.0)[0],
)
await unweighted.async_pre_call_hook(
user_api_key_dict=key, cache=unweighted_cache, data=_ptu_request(), call_type="acompletion"
)
weighted_cache = DualCache()
weighted = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(weighted_cache),
ptu_team_ceiling_resolver=_ptu_ceiling_for("t", "test-model", tpm_limit=1000, ratio=4.0)[0],
)
with pytest.raises(HTTPException) as rejected:
await weighted.async_pre_call_hook(
user_api_key_dict=key, cache=weighted_cache, data=_ptu_request(), call_type="acompletion"
)
assert rejected.value.status_code == 429
assert "model_per_team_ptu" in str(rejected.value.detail)
@pytest.mark.asyncio
async def test_the_ptu_counter_holds_the_normalized_reservation_beside_the_raw_one():
cache = DualCache()
resolve, _ = _ptu_ceiling_for("t", "test-model", tpm_limit=2000, ratio=4.0)
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")
stash = get_request_stash()
assert stash is not None
assert stash.ptu_reserved_tokens == stash.reserved_tokens + 3 * 300
ptu_key = handler.create_rate_limit_keys("model_per_team_ptu", "t:test-model", "tokens")
assert int(await cache.async_get_cache(key=ptu_key) or 0) == stash.ptu_reserved_tokens
@pytest.mark.asyncio
@ -7386,7 +7434,7 @@ 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)
"current": PTUTeamCeiling(tpm_limit=2000, output_to_input_ratio=4.0, cached_input_ratio=0.0)
}
cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
@ -7399,7 +7447,7 @@ async def test_a_reservation_is_settled_even_after_the_teams_share_is_gone():
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
assert stash.ptu_reserved_tokens > stash.reserved_tokens > 150
ceiling["current"] = None
ops = handler._build_success_event_pipeline_operations(
@ -7408,7 +7456,61 @@ async def test_a_reservation_is_settled_even_after_the_teams_share_is_gone():
rate_limit_type="total",
)
assert _ptu_increment(handler, ops) == 150 - stash.reserved_tokens
assert _ptu_increment(handler, ops) == 300 - stash.ptu_reserved_tokens
async def _reserve_a_ptu_minute(cache: DualCache, call_id: str) -> tuple[_PROXY_MaxParallelRequestsHandler, str]:
resolve, _ = _ptu_ceiling_for("t", "test-model", tpm_limit=2000, ratio=4.0)
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(), "litellm_call_id": call_id},
call_type="acompletion",
)
stash = get_request_stash()
assert stash is not None and stash.ptu_reserved_tokens > 0
return handler, handler.create_rate_limit_keys("model_per_team_ptu", "t:test-model", "tokens")
@pytest.mark.asyncio
async def test_a_failed_stream_settles_the_ptu_counter_in_normalized_tokens():
"""The partial usage a failed stream recovered is 20 input and 7 output tokens: 48 normalized
at 4:1, which is what stays in the window instead of the raw 27 or the whole reservation."""
cache = DualCache()
handler, ptu_key = await _reserve_a_ptu_minute(cache, "ptu-partial")
await handler.async_log_failure_event(
kwargs={
"litellm_call_id": "ptu-partial",
"standard_logging_object": {
"metadata": {"user_api_key_hash": hash_token("sk-ptu"), "user_api_key_team_id": "t"}
},
"combined_usage_object": Usage(prompt_tokens=20, completion_tokens=7, total_tokens=27),
},
response_obj=None,
start_time=None,
end_time=None,
)
assert int(await cache.async_get_cache(key=ptu_key) or 0) == 20 + 4 * 7
@pytest.mark.asyncio
async def test_a_proxy_side_rejection_refunds_the_whole_normalized_ptu_reservation():
cache = DualCache()
handler, ptu_key = await _reserve_a_ptu_minute(cache, "ptu-rejected")
await handler.async_post_call_failure_hook(
request_data={**_ptu_request(), "litellm_call_id": "ptu-rejected"},
original_exception=Exception("guardrail rejected the request"),
user_api_key_dict=UserAPIKeyAuth(api_key=hash_token("sk-ptu"), team_id="t"),
)
assert int(await cache.async_get_cache(key=ptu_key) or 0) == 0
def test_usage_that_only_reports_a_total_charges_that_total_to_the_ptu_counter():

View file

@ -4,7 +4,7 @@ from typing import Final
import pytest
from litellm.litellm_core_utils.azure_ptu_capacity import PTUCapacity
from litellm.llms.azure.ptu_capacity import PTUCapacity
from litellm.proxy.management_endpoints.ptu_consumption import attach_ptu_hours
from litellm.types.proxy.management_endpoints.common_daily_activity import (
BreakdownMetrics,

View file

@ -15,7 +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.llms.azure.ptu_capacity import AZURE_PTU_CAPACITY
from litellm.proxy._types import (
LiteLLM_BudgetTable,
LiteLLM_BudgetTableFull,

View file

@ -2,7 +2,7 @@
from typing import Final
from litellm.litellm_core_utils.azure_ptu_capacity import AZURE_PTU_CAPACITY
from litellm.llms.azure.ptu_capacity import AZURE_PTU_CAPACITY
from litellm.router_utils.ptu_shares import (
PTUTeamCeiling,
filter_ptu_shared_deployments,
@ -144,3 +144,24 @@ def test_a_sized_reservation_and_an_unreserved_deployment_raise_no_warning():
assert ptu_capacity_warning("gpt-4.1-ptu", _single_team()) 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
def test_an_unsized_single_team_azure_reservation_is_warned_about_its_ptu_hours_only():
warning: Final = ptu_capacity_warning("gpt-4.1-ptu", _single_team(model="azure/my-ptu-deployment"))
assert warning is not None
assert "PTU hours" in warning
assert "ceiling" not in warning
def test_an_unsized_single_team_reservation_on_another_provider_is_not_warned_about():
assert ptu_capacity_warning("claude-ptu", _single_team(model="anthropic/claude-sonnet-4-5")) is None
def test_a_bare_model_name_counts_as_azure_through_custom_llm_provider():
deployment: Final = {
**_single_team(model="my-ptu-deployment"),
"litellm_params": {"model": "my-ptu-deployment", "custom_llm_provider": "azure"},
}
warning: Final = ptu_capacity_warning("gpt-4.1-ptu", deployment)
assert warning is not None
assert "PTU hours" in warning

View file

@ -17915,14 +17915,17 @@ def test_ptu_shares_raise_when_only_shared_deployments_remain(monkeypatch):
assert [d["model_info"]["id"] for d in deployments] == ["shared-deployment"]
def test_ptu_shares_do_not_filter_while_the_feature_is_off(monkeypatch):
def test_ptu_shares_hide_the_shared_deployment_even_while_the_feature_is_off(monkeypatch):
"""The flag switches cost attribution on; a declared split is an access rule and holds
without it, so a team never reaches a deployment reserved for others while the flag is
off."""
monkeypatch.delenv("LITELLM_ENABLE_PTU_COST_ATTRIBUTION", raising=False)
router = Router(model_list=_shared_ptu_model_list())
_, deployments = router._common_checks_available_deployment(
model="gpt-4.1-ptu",
request_kwargs={"metadata": {"user_api_key_team_id": "team-c"}},
)
assert {d["model_info"]["id"] for d in deployments} == {"shared-deployment", "open-deployment"}
assert [d["model_info"]["id"] for d in deployments] == ["open-deployment"]
def test_a_shared_ptu_deployment_whose_shares_do_not_add_up_is_refused_at_registration(monkeypatch):

View file

@ -2142,7 +2142,7 @@ def test_an_incomplete_reservation_is_refused_rather_than_served(dropped):
@pytest.mark.parametrize(
"dropped, expected",
[
("team_id", "team_id is required when PTU fields are set (one model maps to one team)"),
("team_id", "team_id or ptu_shares is required when PTU fields are set"),
("cost_per_ptu_per_hour", "ptu_count and cost_per_ptu_per_hour must be set together"),
],
ids=["no team_id", "count without rate"],