diff --git a/litellm/litellm_core_utils/azure_ptu_capacity.py b/litellm/llms/azure/ptu_capacity.py similarity index 100% rename from litellm/litellm_core_utils/azure_ptu_capacity.py rename to litellm/llms/azure/ptu_capacity.py diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 4fc994c5939..c61b35f5936 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -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 () diff --git a/litellm/proxy/management_endpoints/ptu_consumption.py b/litellm/proxy/management_endpoints/ptu_consumption.py index 12736b312a8..2e125864c93 100644 --- a/litellm/proxy/management_endpoints/ptu_consumption.py +++ b/litellm/proxy/management_endpoints/ptu_consumption.py @@ -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, diff --git a/litellm/router.py b/litellm/router.py index 79e9f0a2d4f..2275578c695 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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( diff --git a/litellm/router_utils/ptu_shares.py b/litellm/router_utils/ptu_shares.py index 1401dd23d2e..a797cd2d109 100644 --- a/litellm/router_utils/ptu_shares.py +++ b/litellm/router_utils/ptu_shares.py @@ -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 diff --git a/tests/test_litellm/litellm_core_utils/test_azure_ptu_capacity.py b/tests/test_litellm/llms/azure/test_azure_ptu_capacity.py similarity index 98% rename from tests/test_litellm/litellm_core_utils/test_azure_ptu_capacity.py rename to tests/test_litellm/llms/azure/test_azure_ptu_capacity.py index 4d96068cd8b..59e26ecb776 100644 --- a/tests/test_litellm/litellm_core_utils/test_azure_ptu_capacity.py +++ b/tests/test_litellm/llms/azure/test_azure_ptu_capacity.py @@ -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, diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 8f44945718d..09a8314dc37 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -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(): diff --git a/tests/test_litellm/proxy/management_endpoints/test_ptu_consumption.py b/tests/test_litellm/proxy/management_endpoints/test_ptu_consumption.py index e9ef4ea2994..432bf858899 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ptu_consumption.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ptu_consumption.py @@ -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, diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index 7e0792cdc2c..87fff1d5d6a 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -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, diff --git a/tests/test_litellm/router_utils/test_ptu_shares.py b/tests/test_litellm/router_utils/test_ptu_shares.py index 110161bfa93..2c54c3a1c8d 100644 --- a/tests/test_litellm/router_utils/test_ptu_shares.py +++ b/tests/test_litellm/router_utils/test_ptu_shares.py @@ -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 diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index ad9f6e9c4aa..a58292d6d22 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -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): diff --git a/tests/test_litellm/test_router_model_cost_isolation.py b/tests/test_litellm/test_router_model_cost_isolation.py index d73f5efa96b..3c42de71cf7 100644 --- a/tests/test_litellm/test_router_model_cost_isolation.py +++ b/tests/test_litellm/test_router_model_cost_isolation.py @@ -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"],