diff --git a/litellm/litellm_core_utils/azure_ptu_capacity.py b/litellm/litellm_core_utils/azure_ptu_capacity.py index 3ba4b9b6abf..c68429dc7b7 100644 --- a/litellm/litellm_core_utils/azure_ptu_capacity.py +++ b/litellm/litellm_core_utils/azure_ptu_capacity.py @@ -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: diff --git a/litellm/litellm_core_utils/ptu_pricing.py b/litellm/litellm_core_utils/ptu_pricing.py index b68b5cc93c2..e9cd57258ab 100644 --- a/litellm/litellm_core_utils/ptu_pricing.py +++ b/litellm/litellm_core_utils/ptu_pricing.py @@ -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 diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index c9f8e27aa5b..4fc994c5939 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -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, diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 34afa030287..b3dca8b127f 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -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) + ), ) diff --git a/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py b/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py index e2c0ccb0327..1ca0abf19ed 100644 --- a/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py +++ b/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py @@ -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), diff --git a/litellm/router_utils/ptu_shares.py b/litellm/router_utils/ptu_shares.py index 04e54395b5a..1401dd23d2e 100644 --- a/litellm/router_utils/ptu_shares.py +++ b/litellm/router_utils/ptu_shares.py @@ -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 + ) ) diff --git a/litellm/types/router.py b/litellm/types/router.py index 5ead0dd7626..7c91b8f099f 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -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 diff --git a/tests/test_litellm/litellm_core_utils/test_azure_ptu_capacity.py b/tests/test_litellm/litellm_core_utils/test_azure_ptu_capacity.py index d07d71f0d82..4d96068cd8b 100644 --- a/tests/test_litellm/litellm_core_utils/test_azure_ptu_capacity.py +++ b/tests/test_litellm/litellm_core_utils/test_azure_ptu_capacity.py @@ -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(): diff --git a/tests/test_litellm/litellm_core_utils/test_ptu_pricing.py b/tests/test_litellm/litellm_core_utils/test_ptu_pricing.py index ef8d7703924..7dda3432e43 100644 --- a/tests/test_litellm/litellm_core_utils/test_ptu_pricing.py +++ b/tests/test_litellm/litellm_core_utils/test_ptu_pricing.py @@ -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") 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 a18b38453b0..8f44945718d 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 @@ -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( diff --git a/tests/test_litellm/proxy/management_endpoints/test_ptu_model_settings.py b/tests/test_litellm/proxy/management_endpoints/test_ptu_model_settings.py index a7c21d74f1e..8ed3df44216 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ptu_model_settings.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ptu_model_settings.py @@ -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"}) 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 5a1b87b7a01..7e0792cdc2c 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -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(): diff --git a/tests/test_litellm/proxy/spend_tracking/test_ptu_feature_flag.py b/tests/test_litellm/proxy/spend_tracking/test_ptu_feature_flag.py index 7f4bd935a2b..23d0b4ed84f 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_ptu_feature_flag.py +++ b/tests/test_litellm/proxy/spend_tracking/test_ptu_feature_flag.py @@ -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 diff --git a/tests/test_litellm/proxy/spend_tracking/test_ptu_flat_cost_rollup.py b/tests/test_litellm/proxy/spend_tracking/test_ptu_flat_cost_rollup.py index 73e3e5f3ec6..2bcb31c4db2 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_ptu_flat_cost_rollup.py +++ b/tests/test_litellm/proxy/spend_tracking/test_ptu_flat_cost_rollup.py @@ -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) diff --git a/tests/test_litellm/router_utils/test_ptu_shares.py b/tests/test_litellm/router_utils/test_ptu_shares.py index 150cfd31b75..110161bfa93 100644 --- a/tests/test_litellm/router_utils/test_ptu_shares.py +++ b/tests/test_litellm/router_utils/test_ptu_shares.py @@ -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