merge: bring litellm_internal_staging into litellm_mcp_lifecycle_e2e

This commit is contained in:
Yuneng Jiang 2026-09-06 11:46:03 +00:00
commit d85023e38c
No known key found for this signature in database
18 changed files with 1003 additions and 104 deletions

View file

@ -25,6 +25,27 @@ class BatchCostUsageResult:
failed_requests: int
_COMPLETED_BATCH_STATUSES: Final = frozenset({"completed", "complete"})
_TERMINAL_BATCH_STATUSES: Final = _COMPLETED_BATCH_STATUSES | frozenset({"failed", "cancelled", "expired"})
def batch_cost_is_final(batch: Batch) -> bool:
"""Whether this retrieve of the batch is the one to account its cost from.
A batch still in flight has nothing to price, and a "completed" batch can report
no output_file_id for a moment before the output populates; pricing either records
$0 under the batch's single spend row and pins it there. Final means a completed
batch whose output file has arrived or whose counts prove no line succeeded, or
any other terminal status (failed, cancelled, expired).
"""
if batch.status not in _TERMINAL_BATCH_STATUSES:
return False
if batch.status not in _COMPLETED_BATCH_STATUSES or batch.output_file_id is not None:
return True
request_counts: Final = batch.request_counts
return request_counts is not None and request_counts.total > 0 and request_counts.completed == 0
async def calculate_batch_cost_and_usage(
file_content_dictionary: list[dict],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"],

View file

@ -36,7 +36,7 @@ from litellm._logging import (
verbose_logger,
)
from litellm._uuid import uuid
from litellm.batches.batch_utils import _handle_completed_batch
from litellm.batches.batch_utils import _handle_completed_batch, batch_cost_is_final
from litellm.caching.caching import DualCache, InMemoryCache
from litellm.caching.caching_handler import LLMCachingHandler
from litellm.constants import (
@ -2899,13 +2899,6 @@ class Logging(LiteLLMLoggingBaseClass):
): # polling job will query these frequently, don't spam db logs
return
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
)
# check if file id is a unified file id
is_base64_unified_file_id: Final = _is_base64_encoded_unified_file_id(result.id)
batch_cost: Final = kwargs.get("batch_cost", None)
batch_usage = kwargs.get("batch_usage", None)
batch_models = kwargs.get("batch_models", None)
@ -2913,9 +2906,7 @@ class Logging(LiteLLMLoggingBaseClass):
batch_failed_requests: Final = kwargs.get("batch_failed_requests", None)
has_explicit_batch_data: Final = all(x is not None for x in (batch_cost, batch_usage, batch_models))
should_compute_batch_data: Final = (
not is_base64_unified_file_id or not has_explicit_batch_data and result.status == "completed"
)
should_compute_batch_data: Final = not has_explicit_batch_data and batch_cost_is_final(result)
if has_explicit_batch_data:
result._hidden_params["response_cost"] = batch_cost
result._hidden_params["batch_models"] = batch_models

View file

@ -17,6 +17,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
from litellm.llms.azure.common_utils import BaseAzureLLM
from litellm.llms.azure_ai.common_utils import is_foundry_model_inference_base
from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj
from litellm.llms.openai.chat.gpt_5_transformation import OpenAIGPT5Config
from litellm.llms.openai.common_utils import drop_params_from_unprocessable_entity_error
from litellm.llms.openai.openai import OpenAIConfig
from litellm.llms.xai.chat.transformation import XAIChatConfig
@ -42,12 +43,37 @@ NON_OPENAI_SPEC_MESSAGE_FIELDS: Final = (
)
class AzureAIGPT5Config(OpenAIGPT5Config):
@classmethod
def _model_map_lookup_name(cls, model: str) -> str:
"""Normalise a Foundry routing name to its cost-map key, when the map has one.
A Foundry deployment and its OpenAI-hosted namesake are different products with
different capabilities, so ``azure_ai/<model>`` is the entry to read whenever the map
carries it. Most gpt-5-family names have no ``azure_ai/`` row, though, and prefixing
those anyway costs them every flag: ``get_llm_provider`` re-resolves an ``azure_ai/``
name to the azure provider when a global AZURE_AI_API_BASE points at an
openai.azure.com host, ``azure/<model>`` is not a key either, so the lookup lands
nowhere and every effort answer degrades to False. A missing key defers to the base
resolver instead.
"""
prefixed: Final = model if model.startswith("azure_ai/") else f"azure_ai/{model}"
return prefixed if prefixed in litellm.model_cost else super()._model_map_lookup_name(model)
azureAIGPT5Config: Final = AzureAIGPT5Config()
class AzureAIStudioConfig(OpenAIConfig):
def get_supported_openai_params(self, model: str) -> list:
model_supports_tool_choice = True # azure ai supports this by default
if not supports_tool_choice(model=f"azure_ai/{model}"):
model_supports_tool_choice = False
supported_params = super().get_supported_openai_params(model)
supported_params = (
azureAIGPT5Config.get_supported_openai_params(model)
if azureAIGPT5Config.is_model_gpt_5_model(model)
else super().get_supported_openai_params(model)
)
if not model_supports_tool_choice:
filtered_supported_params: Final = []
for param in supported_params:
@ -61,6 +87,27 @@ class AzureAIStudioConfig(OpenAIConfig):
return supported_params
def map_openai_params(
self,
non_default_params: dict[str, object], # mutable-ok: OpenAIConfig.map_openai_params signature
optional_params: dict[str, object], # mutable-ok: OpenAIConfig.map_openai_params signature
model: str,
drop_params: bool,
) -> dict[str, object]: # mutable-ok: OpenAIConfig.map_openai_params signature
if not azureAIGPT5Config.is_model_gpt_5_model(model):
return super().map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
model=model,
drop_params=drop_params,
)
return azureAIGPT5Config.map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
model=model,
drop_params=drop_params,
)
def _supports_stop_reason(self, model: str) -> bool:
"""
Check if the model supports stop tokens.

View file

@ -3485,6 +3485,55 @@
"supports_response_schema": true,
"supports_tool_choice": true
},
"azure_ai/gpt-6-astra": {
"cache_creation_input_token_cost": 1.25e-05,
"cache_creation_input_token_cost_above_272k_tokens": 2.5e-05,
"cache_read_input_token_cost": 1e-06,
"cache_read_input_token_cost_above_272k_tokens": 2e-06,
"input_cost_per_token": 1e-05,
"input_cost_per_token_above_272k_tokens": 2e-05,
"litellm_provider": "azure_ai",
"max_input_tokens": 922000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 5e-05,
"output_cost_per_token_above_272k_tokens": 7.5e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"source": "https://ai.azure.com/catalog/models/gpt-6-astra",
"supported_endpoints": [
"/v1/chat/completions",
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_computer_use": true,
"supports_function_calling": true,
"supports_max_reasoning_effort": false,
"supports_minimal_reasoning_effort": false,
"supports_native_streaming": true,
"supports_none_reasoning_effort": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_cache_breakpoint": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": true,
"supports_xhigh_reasoning_effort": true
},
"azure_ai/gpt-5.5": {
"deprecation_date": "2027-10-26",
"cache_read_input_token_cost": 5e-07,
@ -7189,7 +7238,7 @@
],
"supports_computer_use": true,
"supports_function_calling": true,
"supports_max_reasoning_effort": true,
"supports_max_reasoning_effort": false,
"supports_minimal_reasoning_effort": false,
"supports_native_streaming": true,
"supports_none_reasoning_effort": true,
@ -7455,7 +7504,7 @@
],
"supports_computer_use": true,
"supports_function_calling": true,
"supports_max_reasoning_effort": true,
"supports_max_reasoning_effort": false,
"supports_minimal_reasoning_effort": false,
"supports_native_streaming": true,
"supports_none_reasoning_effort": true,

View file

@ -14,6 +14,7 @@ import time
import traceback
from collections.abc import Mapping, Sequence
from datetime import datetime, timedelta, timezone
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, cast, overload
import litellm
@ -84,6 +85,25 @@ else:
RESPONSES_SESSION_CALL_TYPES: Final = frozenset({CallTypes.responses.value, CallTypes.aresponses.value})
def _is_batch_cost_row(payload: SpendLogsPayload) -> bool:
return payload.get("call_type") == CallTypes.aretrieve_batch.value and payload.get("status") == "success"
_BATCH_COST_CLAIM_FIELDS: Final = frozenset({"request_id", "call_type", "spend", "startTime", "endTime", "status"})
def _batch_cost_row_to_write(payload: SpendLogsPayload, disable_spend_logs: bool) -> Mapping[str, object]:
"""Reduce a batch's cost row to what tells the retrieves apart when logging is off.
A proxy run with spend logs disabled still needs one row per batch to charge it once,
so the row is written either way, but it carries no request of its own: no metadata,
no requester IP, no key, model, or token counts (LIT-7048).
"""
if disable_spend_logs is False:
return payload
return MappingProxyType({field: value for field, value in payload.items() if field in _BATCH_COST_CLAIM_FIELDS})
class _SpendBatch(Protocol):
litellm_usertable: BatchTable
litellm_verificationtoken: BatchTable
@ -215,7 +235,12 @@ class DBSpendUpdateWriter:
start_time: datetime | None,
end_time: datetime | None,
response_cost: float | None,
) -> None:
) -> bool:
"""Record the request's spend, answering whether its cost still needs charging.
False only for a batch retrieve whose cost row another retrieve already wrote,
so the caller leaves the key, team, and user counters alone (LIT-7048).
"""
from litellm.proxy.proxy_server import (
disable_spend_logs,
litellm_proxy_budget_name,
@ -232,7 +257,7 @@ class DBSpendUpdateWriter:
team_id,
)
if ProxyUpdateSpend.disable_spend_updates() is True:
return
return True
if token is not None and isinstance(token, str) and token.startswith("sk-"):
hashed_token = hash_token(token=token)
else:
@ -262,11 +287,12 @@ class DBSpendUpdateWriter:
if team_id is not None and team_id != "":
payload["team_id"] = team_id
if not await self._record_spend_log(
payload=payload, prisma_client=prisma_client, disable_spend_logs=disable_spend_logs
):
return False
if disable_spend_logs is False:
await self._insert_spend_log_to_db(
payload=payload,
prisma_client=prisma_client,
)
await self._enqueue_tool_usage_transaction(
payload=payload,
completion_response=completion_response,
@ -306,6 +332,7 @@ class DBSpendUpdateWriter:
)
verbose_proxy_logger.debug("Runs spend update on all tables")
return True
except Exception:
spend_log_error(
"Spend tracking - update_database failed. Spend log insertion or daily transaction enqueue "
@ -318,7 +345,102 @@ class DBSpendUpdateWriter:
org_id,
end_user_id,
)
return
return True
async def _record_spend_log(
self, payload: SpendLogsPayload, prisma_client: "PrismaClient | None", disable_spend_logs: bool
) -> bool:
if prisma_client is not None and _is_batch_cost_row(payload):
return await self._claim_batch_cost_spend_log(
payload=payload, prisma_client=prisma_client, disable_spend_logs=disable_spend_logs
)
if disable_spend_logs is False:
await self._insert_spend_log_to_db(payload=payload, prisma_client=prisma_client)
return True
async def _claim_batch_cost_spend_log(
self, payload: SpendLogsPayload, prisma_client: "PrismaClient", disable_spend_logs: bool
) -> bool:
"""Write the batch's cost row now, or learn that another retrieve already did.
Every retrieve of one batch shares this row, so the insert that lands first owns
the charge and every later one finds the row and charges nothing (LIT-7048). Only
a row that recorded a charge counts: a failed retrieve, a request whose client
picked the batch id as its call id, and the $0 row an older proxy left behind
while the batch was still running all leave the charge to be made.
"""
from litellm.repositories.table_repositories import SpendLogsRepository
request_id: Final = payload["request_id"]
row: Final = _batch_cost_row_to_write(payload, disable_spend_logs)
spend_logs: Final = SpendLogsRepository(prisma_client).table
try:
claimed: Final = await spend_logs.create_many(
data=[prisma_client.jsonify_object(row)], # mutable-ok: prisma create_many takes a list
skip_duplicates=True,
)
if claimed == 1:
return True
existing: Final = await spend_logs.find_unique(
where={"request_id": request_id} # mutable-ok: prisma where clause
)
except Exception as e: # noqa: BLE001 # prisma raises its own hierarchy; an unreachable DB queues the row like any other spend log
verbose_proxy_logger.warning(
"Could not claim spend row %s for a batch's cost, queueing it: %s", request_id, e
)
await self._insert_spend_log_to_db(payload=prisma_client.jsonify_object(row), prisma_client=prisma_client)
return True
if existing is None or existing.call_type != CallTypes.aretrieve_batch.value or existing.status != "success":
verbose_proxy_logger.warning(
"Spend row %s belongs to a %s request, so this batch's cost is charged without a row of its own",
request_id,
getattr(existing, "call_type", None),
)
return True
if existing.spend > 0:
verbose_proxy_logger.debug("Cost tracking skipped: spend row %s already charged this batch", request_id)
return False
return await self._take_over_uncharged_batch_cost_row(payload=payload, prisma_client=prisma_client, row=row)
async def _take_over_uncharged_batch_cost_row(
self, payload: SpendLogsPayload, prisma_client: "PrismaClient", row: Mapping[str, object]
) -> bool:
"""Take the batch's cost row over from the poll that left it charging nothing.
A pre-upgrade proxy wrote that row every time it polled the batch while it was still
running, so the charge is still to be made and the row still has to end up carrying
it. The row stops matching the moment it carries a charge, so it is one retrieve that
takes it over and charges, and every later one reads the charge and charges nothing.
"""
from litellm.repositories.table_repositories import SpendLogsRepository
request_id: Final = payload["request_id"]
if payload["spend"] <= 0:
verbose_proxy_logger.debug(
"Cost tracking skipped: this batch costs nothing and spend row %s says so", request_id
)
return False
try:
taken_over: Final = await SpendLogsRepository(prisma_client).table.update_many(
data=prisma_client.jsonify_object(
MappingProxyType({field: value for field, value in row.items() if field != "request_id"})
),
where={ # mutable-ok: prisma where clause
"request_id": request_id,
"call_type": CallTypes.aretrieve_batch.value,
"status": "success",
"spend": 0.0,
},
)
except Exception as e: # noqa: BLE001 # prisma raises its own hierarchy; the next retrieve takes the row over
verbose_proxy_logger.warning(
"Could not take over spend row %s, leaving this batch's cost to the next retrieve: %s", request_id, e
)
return False
if taken_over == 0:
verbose_proxy_logger.debug("Cost tracking skipped: spend row %s already charged this batch", request_id)
return False
return True
async def _enqueue_tool_usage_transaction(
self,

View file

@ -6,6 +6,7 @@ from typing import TYPE_CHECKING, Any, Final, cast
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.batches.batch_utils import batch_cost_is_final
from litellm.constants import BACKGROUND_INTERACTION_COST_POLLING_ENABLED
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import (
@ -37,6 +38,7 @@ from litellm.proxy.spend_tracking.spend_tracking_utils import (
from litellm.proxy.utils import ProxyUpdateSpend
from litellm.types.utils import (
CallTypes,
LiteLLMBatch,
StandardLoggingPayload,
StandardLoggingPayloadErrorInformation,
)
@ -248,6 +250,18 @@ class _ProxyDBLogger(CustomLogger):
)
_write_spend_metadata_to_kwargs(kwargs=kwargs, metadata=metadata)
budget_reservation: Final = _get_budget_reservation_from_metadata(metadata=metadata)
if (
isinstance(completion_response, LiteLLMBatch)
and kwargs.get("call_type") == CallTypes.aretrieve_batch.value
and not batch_cost_is_final(completion_response)
):
verbose_proxy_logger.debug(
"Cost tracking deferred for batch %s still in status %s",
completion_response.id,
completion_response.status,
)
await _release_budget_reservation(budget_reservation=budget_reservation)
return
user_id: Final = cast(str | None, metadata.get("user_api_key_user_id", None))
team_id: Final = cast(str | None, metadata.get("user_api_key_team_id", None))
org_id: Final = cast(str | None, metadata.get("user_api_key_org_id", None))
@ -285,7 +299,7 @@ class _ProxyDBLogger(CustomLogger):
call_type=call_type,
):
## UPDATE DATABASE
await _update_database_and_spend_counters(
charged: Final = await _update_database_and_spend_counters(
proxy_logging_obj=proxy_logging_obj,
increment_spend_counters=increment_spend_counters,
user_api_key=user_api_key,
@ -302,6 +316,8 @@ class _ProxyDBLogger(CustomLogger):
request_tags=tags,
model_access_groups=model_access_groups,
)
if not charged:
return
# update cache (fire-and-forget for backward compat:
# cached object fields, soft budget alerts, etc.)
@ -578,9 +594,9 @@ async def _update_database_and_spend_counters(
budget_reservation: dict | None,
request_tags: list[str] | None = None,
model_access_groups: Sequence[str] | None = None,
) -> None:
) -> bool:
try:
await proxy_logging_obj.db_spend_update_writer.update_database(
charged: Final = await proxy_logging_obj.db_spend_update_writer.update_database(
token=user_api_key,
response_cost=response_cost,
user_id=user_id,
@ -605,6 +621,9 @@ async def _update_database_and_spend_counters(
"Failed to invalidate budget reservation counters after release failed"
)
raise
if not charged:
await _release_budget_reservation(budget_reservation=budget_reservation)
return False
try:
await increment_spend_counters(
@ -630,6 +649,7 @@ async def _update_database_and_spend_counters(
finally:
budget_reservation["finalized"] = True
raise
return True
async def _release_budget_reservation(budget_reservation: dict | None) -> None:

View file

@ -15,6 +15,7 @@ from typing import (
runtime_checkable,
)
from litellm.batches.batch_utils import batch_cost_is_final
from litellm.proxy._types import ProxyException
from litellm.repositories.table_repositories import (
ManagedFileRepository,
@ -1357,12 +1358,7 @@ def _completed_batch_safe_to_retire(response: "LiteLLMBatch") -> bool:
enumerated the batch and none succeeded. A zero or unknown total means counts
are unreported, so stay eligible and let the next poller pass revisit it. (#37713)
"""
if response.output_file_id is not None:
return True
request_counts = response.request_counts
if request_counts is None:
return False
return request_counts.total > 0 and request_counts.completed == 0
return batch_cost_is_final(response)
async def update_batch_in_database(

View file

@ -10,8 +10,8 @@ opt-in. none is opt-out everywhere except the azure gpt-5 family, whose config r
UnsupportedParamsError without an explicit true.
xhigh is gated on the request path by the openai and azure gpt-5 configs. max is not gated there at
all: every entry carrying supports_max_reasoning_effort is Claude-family, and
anthropic/chat/transformation.py gates max on the output_config path while its reasoning_effort
all: outside the gpt-6-astra rows every entry carrying supports_max_reasoning_effort is Claude-family,
and anthropic/chat/transformation.py gates max on the output_config path while its reasoning_effort
path maps any level to a thinking budget. Making max opt-in is a deliberate trade, then, since an
explicit flag is the only signal that the tier is a real one rather than litellm rounding the level
to a budget, and a missing flag costs advisory metadata rather than a rejected request.

View file

@ -3485,6 +3485,55 @@
"supports_response_schema": true,
"supports_tool_choice": true
},
"azure_ai/gpt-6-astra": {
"cache_creation_input_token_cost": 1.25e-05,
"cache_creation_input_token_cost_above_272k_tokens": 2.5e-05,
"cache_read_input_token_cost": 1e-06,
"cache_read_input_token_cost_above_272k_tokens": 2e-06,
"input_cost_per_token": 1e-05,
"input_cost_per_token_above_272k_tokens": 2e-05,
"litellm_provider": "azure_ai",
"max_input_tokens": 922000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 5e-05,
"output_cost_per_token_above_272k_tokens": 7.5e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"source": "https://ai.azure.com/catalog/models/gpt-6-astra",
"supported_endpoints": [
"/v1/chat/completions",
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_computer_use": true,
"supports_function_calling": true,
"supports_max_reasoning_effort": false,
"supports_minimal_reasoning_effort": false,
"supports_native_streaming": true,
"supports_none_reasoning_effort": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_cache_breakpoint": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": true,
"supports_xhigh_reasoning_effort": true
},
"azure_ai/gpt-5.5": {
"deprecation_date": "2027-10-26",
"cache_read_input_token_cost": 5e-07,
@ -7189,7 +7238,7 @@
],
"supports_computer_use": true,
"supports_function_calling": true,
"supports_max_reasoning_effort": true,
"supports_max_reasoning_effort": false,
"supports_minimal_reasoning_effort": false,
"supports_native_streaming": true,
"supports_none_reasoning_effort": true,
@ -7455,7 +7504,7 @@
],
"supports_computer_use": true,
"supports_function_calling": true,
"supports_max_reasoning_effort": true,
"supports_max_reasoning_effort": false,
"supports_minimal_reasoning_effort": false,
"supports_native_streaming": true,
"supports_none_reasoning_effort": true,

View file

@ -21,11 +21,12 @@ from types import MappingProxyType
import httpx
import pytest
import respx
from openai.types.batch import BatchRequestCounts
import litellm
import litellm.batches.batch_utils as bu
from litellm.types.utils import Usage
from litellm.types.utils import LiteLLMBatch, Usage
# --------------------------------------------------------------------------- #
# Builders for batch OUTPUT file rows.
@ -1718,3 +1719,57 @@ def test_unparsable_bedrock_batch_usage_warns(caplog):
assert usage.total_tokens == 0
assert "does not understand" in caplog.text
assert "inputTextTokenCount" in caplog.text
# --------------------------------------------------------------------------- #
# batch_cost_is_final
# --------------------------------------------------------------------------- #
def _retrieved_batch(
status: str, output_file_id: str | None = None, counts: BatchRequestCounts | None = None
) -> LiteLLMBatch:
return LiteLLMBatch(
id="batch_abc",
completion_window="24h",
created_at=1,
endpoint="/v1/chat/completions",
input_file_id="file-in",
object="batch",
status="validating",
output_file_id=output_file_id,
request_counts=counts,
).model_copy(update={"status": status})
class TestBatchCostIsFinal:
"""Every retrieve of one batch writes the same spend row, so the first retrieve
that prices it decides the row for good. A poll before the output exists must
therefore not count as final: pricing it recorded $0 and pinned it (LIT-7048)."""
@pytest.mark.parametrize("status", ["validating", "in_progress", "finalizing", "cancelling"])
def test_in_flight_batch_is_not_final(self, status):
assert bu.batch_cost_is_final(_retrieved_batch(status)) is False
@pytest.mark.parametrize("status", ["completed", "complete"])
def test_completed_with_output_is_final(self, status):
assert bu.batch_cost_is_final(_retrieved_batch(status, output_file_id="file-out")) is True
def test_completed_without_output_and_unknown_counts_is_not_final(self):
assert bu.batch_cost_is_final(_retrieved_batch("completed")) is False
def test_completed_without_output_and_zero_counts_is_not_final(self):
counts = BatchRequestCounts(total=0, completed=0, failed=0)
assert bu.batch_cost_is_final(_retrieved_batch("completed", counts=counts)) is False
def test_completed_without_output_but_successful_lines_is_not_final(self):
counts = BatchRequestCounts(total=2, completed=2, failed=0)
assert bu.batch_cost_is_final(_retrieved_batch("completed", counts=counts)) is False
@pytest.mark.parametrize("status", ["completed", "complete"])
def test_completed_without_output_and_every_line_failed_is_final(self, status):
counts = BatchRequestCounts(total=2, completed=0, failed=2)
assert bu.batch_cost_is_final(_retrieved_batch(status, counts=counts)) is True
@pytest.mark.parametrize("status", ["failed", "expired", "cancelled"])
def test_other_terminal_statuses_are_final(self, status):
assert bu.batch_cost_is_final(_retrieved_batch(status)) is True

View file

@ -2008,7 +2008,14 @@ def test_generic_cost_per_token_azure_gpt56(_local_model_cost_map,
assert round(completion_cost, 10) == round(output_cost * completion_tokens, 10)
@pytest.mark.parametrize("model,zone_multiplier", [("azure/gpt-6-astra", 1.0), ("azure/us/gpt-6-astra", 1.1)])
@pytest.mark.parametrize(
"model,custom_llm_provider,zone_multiplier",
[
("azure/gpt-6-astra", "azure", 1.0),
("azure/us/gpt-6-astra", "azure", 1.1),
("azure_ai/gpt-6-astra", "azure_ai", 1.0),
],
)
@pytest.mark.parametrize(
"prompt_tokens,input_side_multiplier,output_multiplier",
[(100000, 1.0, 1.0), (300000, 2.0, 1.5)],
@ -2016,6 +2023,7 @@ def test_generic_cost_per_token_azure_gpt56(_local_model_cost_map,
def test_generic_cost_per_token_azure_gpt_6_astra_foundry_price_sheet(
_local_model_cost_map,
model,
custom_llm_provider,
zone_multiplier,
prompt_tokens,
input_side_multiplier,
@ -2023,7 +2031,8 @@ def test_generic_cost_per_token_azure_gpt_6_astra_foundry_price_sheet(
):
"""Microsoft Foundry sells gpt-6-astra at the OpenAI rates: $10 input, $1 cache read, $12.50 cache write,
$50 output per 1M tokens on Standard Global, with the input side doubling and output 1.5x above 272K
prompt tokens. Standard US Data Zone carries the usual 10% uplift on every rate.
prompt tokens. Standard US Data Zone carries the usual 10% uplift on every rate. A Foundry
deployment reached through the azure_ai route bills the same Standard Global sheet.
"""
cached_tokens = 50000
cache_write_tokens = 40000
@ -2041,7 +2050,7 @@ def test_generic_cost_per_token_azure_gpt_6_astra_foundry_price_sheet(
prompt_cost, completion_cost = generic_cost_per_token(
model=model,
usage=usage,
custom_llm_provider="azure",
custom_llm_provider=custom_llm_provider,
)
input_side = zone_multiplier * input_side_multiplier
@ -2051,6 +2060,18 @@ def test_generic_cost_per_token_azure_gpt_6_astra_foundry_price_sheet(
assert completion_cost == pytest.approx(zone_multiplier * output_multiplier * completion_tokens * 5e-5)
def test_generic_cost_per_token_azure_ai_gpt_6_astra_flex_bills_the_standard_rate(_local_model_cost_map):
usage = Usage(prompt_tokens=1000, completion_tokens=100, total_tokens=1100)
standard = generic_cost_per_token(model="azure_ai/gpt-6-astra", usage=usage, custom_llm_provider="azure_ai")
flex = generic_cost_per_token(
model="azure_ai/gpt-6-astra", usage=usage, custom_llm_provider="azure_ai", service_tier="flex"
)
assert flex == standard
assert standard == pytest.approx((1000 * 1e-05, 100 * 5e-05))
@pytest.mark.parametrize(
"model,expected_none,expected_xhigh,expected_minimal",
[

View file

@ -632,6 +632,86 @@ class TestRetrieveBatchCostPassesModelIdentity:
assert captured["model_info"]["input_cost_per_token"] == 0.0
class TestRetrieveBatchPricesOnlyFinalBatches:
"""Regression (LIT-7048): retrieving a provider-id batch priced it on every poll.
Every retrieve of one batch logs under the same spend row, so pricing a poll
that landed before the output existed wrote that row at $0 and pinned it there.
Only a final batch gets priced; an in-flight poll carries no cost at all.
"""
@staticmethod
def _logging_obj() -> LitellmLogging:
obj = LitellmLogging(
model="gpt-5.6-luna",
messages=[{"role": "user", "content": "Hey"}],
stream=False,
call_type="aretrieve_batch",
start_time=time.time(),
litellm_call_id="batch-call-2",
function_id="f",
)
obj.custom_llm_provider = "openai"
return obj
@staticmethod
def _batch(status: str, output_file_id: str | None):
from litellm.types.utils import LiteLLMBatch
return LiteLLMBatch(
id="batch_6a9c99e185588190877d391f8b9d7f8a",
completion_window="24h",
created_at=1,
endpoint="/v1/chat/completions",
input_file_id="file-in",
object="batch",
status="validating",
output_file_id=output_file_id,
).model_copy(update={"status": status})
@pytest.mark.asyncio
@pytest.mark.parametrize(
("status", "output_file_id"),
[("validating", None), ("in_progress", None), ("finalizing", None), ("completed", None), ("complete", None)],
)
async def test_non_final_batch_is_not_priced(self, monkeypatch, status, output_file_id) -> None:
from litellm.litellm_core_utils import litellm_logging as logging_module
handle_completed_batch = AsyncMock()
monkeypatch.setattr(logging_module, "_handle_completed_batch", handle_completed_batch)
batch = self._batch(status, output_file_id)
await self._logging_obj()._async_success_handler_body(result=batch, start_time=None, end_time=None)
handle_completed_batch.assert_not_awaited()
assert "response_cost" not in batch._hidden_params
@pytest.mark.asyncio
async def test_completed_batch_with_output_is_priced(self, monkeypatch) -> None:
from litellm.batches.batch_utils import BatchCostUsageResult
from litellm.litellm_core_utils import litellm_logging as logging_module
from litellm.types.utils import Usage
handle_completed_batch = AsyncMock(
return_value=BatchCostUsageResult(
cost=8e-06,
usage=Usage(prompt_tokens=26, completion_tokens=9, total_tokens=35),
models=["gpt-5.6-luna"],
successful_requests=2,
failed_requests=0,
)
)
monkeypatch.setattr(logging_module, "_handle_completed_batch", handle_completed_batch)
batch = self._batch("completed", "file-out")
await self._logging_obj()._async_success_handler_body(result=batch, start_time=None, end_time=None)
handle_completed_batch.assert_awaited_once()
assert batch._hidden_params["response_cost"] == 8e-06
assert batch.usage is not None
assert batch.usage.total_tokens == 35
class TestAnthropicPassthroughCustomPricing:
"""Verify the Anthropic pass-through handler forwards custom pricing."""

View file

@ -82,3 +82,22 @@ class TestTheNormalizedTierIsTheTierSent:
self, local_model_cost_map, model, provider, effort, expected
):
assert _reasoning_effort_sent(model, provider, effort) == expected
@pytest.mark.parametrize(
"model, provider",
[
("gpt-6-astra", "azure_ai"),
("azure_ai/gpt-6-astra", "azure_ai"),
("gpt-6-astra", "azure"),
("us/gpt-6-astra", "azure"),
],
)
def test_an_azure_hosted_astra_deployment_drops_to_the_tier_it_accepts(
self, local_model_cost_map, model, provider
):
"""The deployment answers ``max`` with a 400 naming ``none`` through ``xhigh``, so the rows
say so and the adapter sends the tier below instead of the rejected one."""
assert _reasoning_effort_sent(model, provider, "max") == "xhigh"
def test_the_openai_hosted_twin_still_sends_max(self, local_model_cost_map):
assert _reasoning_effort_sent("gpt-6-astra", "openai", "max") == "max"

View file

@ -3,6 +3,8 @@ from unittest.mock import MagicMock, patch
import pytest
import litellm
from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map
from litellm.llms.azure_ai.azure_model_router.transformation import (
AzureModelRouterConfig,
)
@ -138,6 +140,46 @@ def test_azure_ai_validate_environment_with_azure_ad_token():
assert headers["Content-Type"] == "application/json"
@pytest.fixture
def _local_model_cost_map(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", get_model_cost_map(url=litellm.model_cost_map_url))
def test_foundry_gpt_6_astra_keeps_sampling_params_when_reasoning_effort_is_none(_local_model_cost_map):
optional_params = AzureAIStudioConfig().map_openai_params(
non_default_params={"reasoning_effort": "none", "temperature": 0.2, "top_p": 0.9},
optional_params={},
model="gpt-6-astra",
drop_params=False,
)
assert optional_params == {"reasoning_effort": "none", "temperature": 0.2, "top_p": 0.9}
def test_a_gpt_5_name_without_a_foundry_row_keeps_reading_its_own_entry(
monkeypatch: pytest.MonkeyPatch, _local_model_cost_map
):
"""Most gpt-5-family names have no azure_ai/ row. Reading an azure_ai/ key for those finds
nothing, and an openai.azure.com base sends the name down the azure provider, which has no key
for it either, so every effort answer would silently fall back to false and take temperature,
top_p and logprobs down with it."""
monkeypatch.setenv("AZURE_AI_API_BASE", "https://example-resource.openai.azure.com")
monkeypatch.setenv("AZURE_AI_API_KEY", "placeholder")
optional_params = litellm.utils.get_optional_params(
model="gpt-5.1-chat-latest",
custom_llm_provider="azure_ai",
temperature=0.2,
top_p=0.9,
logprobs=True,
)
assert optional_params["temperature"] == 0.2
assert optional_params["top_p"] == 0.9
assert optional_params["logprobs"] is True
def test_azure_ai_grok_stop_parameter_handling():
"""
Test that Grok models properly handle stop parameter filtering in Azure AI Studio.

View file

@ -857,6 +857,24 @@ def test_add_known_models_refreshes_models_by_provider_for_wildcard_expansion():
litellm.add_known_models(model_cost_map={})
assert fake_model not in litellm.models_by_provider["vertex_ai"]
def test_azure_ai_wildcard_lists_the_foundry_gpt_6_astra_entry(monkeypatch):
import litellm
from litellm.proxy.auth.model_checks import get_known_models_from_wildcard
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
foundry_key = "azure_ai/gpt-6-astra"
local_entry = litellm.get_model_cost_map(url="")[foundry_key]
registered_before = foundry_key in litellm.azure_ai_models
try:
litellm.add_known_models(model_cost_map={foundry_key: local_entry})
assert foundry_key in get_known_models_from_wildcard("azure_ai/*")
finally:
if not registered_before:
litellm.azure_ai_models.discard(foundry_key)
litellm.add_known_models(model_cost_map={})
def test_get_complete_model_list_drops_no_default_models_sentinel():
from litellm.proxy.auth.model_checks import get_complete_model_list

View file

@ -7,6 +7,7 @@ import re
from collections.abc import Callable
from contextlib import asynccontextmanager
from datetime import datetime, timezone
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, call, patch
import pytest
@ -2936,7 +2937,9 @@ async def test_commit_spend_updates_retries_deadlock_on_every_entity_path(monkey
"call_type, expects_flush",
[("aresponses", True), ("responses", True), ("acompletion", False)],
)
async def test_insert_spend_log_asks_for_an_immediate_flush_on_responses_calls(call_type: str, expects_flush: bool):
async def test_insert_spend_log_asks_for_an_immediate_flush_on_rows_other_workers_read_back(
call_type: str, expects_flush: bool
):
"""
A `previous_response_id` chained straight off the previous turn reads the DB, so a
Responses row cannot sit in this worker's queue until the monitor's next poll.
@ -2957,6 +2960,303 @@ async def test_insert_spend_log_asks_for_an_immediate_flush_on_responses_calls(c
PrismaClient.spend_log_flush_requested.clear()
def _batch_cost_payload() -> dict:
return {
**_minimal_spend_payload(),
"request_id": "batch_abc_batch_cost",
"call_type": "aretrieve_batch",
"status": "success",
}
def _spend_logs_prisma(inserted: int, existing: object, taken_over: int = 1) -> MagicMock:
prisma = _tool_usage_prisma()
prisma.jsonify_object = lambda data: dict(data)
prisma.db.litellm_spendlogs.create_many = AsyncMock(return_value=inserted)
prisma.db.litellm_spendlogs.find_unique = AsyncMock(return_value=existing)
prisma.db.litellm_spendlogs.update_many = AsyncMock(return_value=taken_over)
return prisma
async def _update_database_with(
db_writer: DBSpendUpdateWriter,
prisma: MagicMock,
payload: dict,
disable_spend_logs: bool = False,
response_cost: float = 0.25,
) -> bool:
with (
patch( # test-quality-ok: update_database reads this proxy_server global at call time, no seam
"litellm.proxy.proxy_server.disable_spend_logs", disable_spend_logs
),
patch( # test-quality-ok: update_database reads this proxy_server global at call time, no seam
"litellm.proxy.proxy_server.prisma_client", prisma
),
patch( # test-quality-ok: update_database reads this proxy_server global at call time, no seam
"litellm.proxy.proxy_server.litellm_proxy_budget_name", "test-budget"
),
patch( # test-quality-ok: update_database imports the payload builder inside its body, no seam
"litellm.proxy.spend_tracking.spend_tracking_utils.get_logging_payload",
return_value=payload,
),
):
charged = await db_writer.update_database(
token="test-token",
user_id="test-user",
end_user_id=None,
team_id=None,
org_id=None,
kwargs={"model": "gpt-5.6-luna", "call_type": "aretrieve_batch"},
completion_response=None,
start_time=datetime.now(timezone.utc),
end_time=datetime.now(timezone.utc),
response_cost=response_cost,
)
await asyncio.sleep(0)
return charged
@pytest.mark.asyncio
@pytest.mark.parametrize(
("inserted", "existing", "charged"),
[
(1, None, True),
(0, SimpleNamespace(call_type="aretrieve_batch", status="success", spend=0.25), False),
(0, SimpleNamespace(call_type="aretrieve_batch", status="success", spend=0.0), True),
(0, SimpleNamespace(call_type="aretrieve_batch", status="failure", spend=0.0), True),
(0, SimpleNamespace(call_type="aembedding", status="success", spend=0.25), True),
(0, None, True),
],
ids=[
"first_retrieve_owns_the_row",
"another_retrieve_already_charged",
"an_older_proxy_left_a_zero_row_while_the_batch_ran",
"failed_retrieve_holds_the_row",
"client_chosen_call_id_holds_the_row",
"row_gone_between_insert_and_lookup",
],
)
async def test_update_database_charges_a_batch_only_from_the_retrieve_that_wrote_its_row(
inserted: int, existing: object, charged: bool
):
"""
Every retrieve of one batch shares one spend row, so the insert that lands first is
the charge and every later retrieve must leave the counters alone (LIT-7048). A row
that recorded no charge must not be able to take the charge away: neither one a
client planted under the batch id, nor the $0 row a pre-upgrade proxy wrote every
time it polled the batch while it was still running.
"""
db_writer = DBSpendUpdateWriter()
db_writer._batch_database_updates = AsyncMock()
prisma = _spend_logs_prisma(inserted, existing)
assert await _update_database_with(db_writer, prisma, _batch_cost_payload()) is charged
claimed_rows = prisma.db.litellm_spendlogs.create_many.await_args.kwargs
assert claimed_rows["skip_duplicates"] is True
assert [(row["request_id"], row["spend"]) for row in claimed_rows["data"]] == [("batch_abc_batch_cost", 0.25)]
assert prisma.spend_log_transactions == []
assert db_writer._batch_database_updates.await_count == (1 if charged else 0)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("taken_over", "charged"),
[(1, True), (0, False)],
ids=["this_retrieve_takes_it_over", "another_one_got_there_first"],
)
async def test_update_database_charges_a_batch_whose_row_a_pre_upgrade_poll_left_at_zero(
taken_over: int, charged: bool
):
"""
A proxy without this fix wrote the batch's row at $0 on every poll of a running batch,
and the row outlives the upgrade, so the charge has to land on the row itself. Charging
without writing it there would charge again on every later retrieve (LIT-7048).
"""
db_writer = DBSpendUpdateWriter()
db_writer._batch_database_updates = AsyncMock()
existing = SimpleNamespace(call_type="aretrieve_batch", status="success", spend=0.0)
prisma = _spend_logs_prisma(0, existing, taken_over)
assert await _update_database_with(db_writer, prisma, _batch_cost_payload()) is charged
taken = prisma.db.litellm_spendlogs.update_many.await_args.kwargs
assert taken["where"] == {
"request_id": "batch_abc_batch_cost",
"call_type": "aretrieve_batch",
"status": "success",
"spend": 0.0,
}
assert taken["data"]["spend"] == 0.25
assert "request_id" not in taken["data"]
assert db_writer._batch_database_updates.await_count == (1 if charged else 0)
@pytest.mark.asyncio
async def test_update_database_leaves_a_batch_whose_zero_row_it_could_not_take_over_to_the_next_retrieve():
"""
A DB that refuses the takeover leaves the row reading $0, so charging here would charge
the batch again on every later retrieve. The retrieve that does take the row over is the
one that charges.
"""
db_writer = DBSpendUpdateWriter()
db_writer._batch_database_updates = AsyncMock()
existing = SimpleNamespace(call_type="aretrieve_batch", status="success", spend=0.0)
prisma = _spend_logs_prisma(0, existing)
prisma.db.litellm_spendlogs.update_many = AsyncMock(side_effect=RuntimeError("db unreachable"))
assert await _update_database_with(db_writer, prisma, _batch_cost_payload()) is False
assert db_writer._batch_database_updates.await_count == 0
@pytest.mark.asyncio
async def test_update_database_leaves_a_batch_that_cost_nothing_to_the_retrieve_that_wrote_its_row():
"""
A batch every line of which failed costs $0, so its row reads $0 for the honest reason
and the retrieve that wrote it is still the one that accounted it. Taking that row over
on every later retrieve would count one batch as many requests.
"""
db_writer = DBSpendUpdateWriter()
db_writer._batch_database_updates = AsyncMock()
existing = SimpleNamespace(call_type="aretrieve_batch", status="success", spend=0.0)
prisma = _spend_logs_prisma(0, existing)
assert await _update_database_with(db_writer, prisma, _batch_cost_payload(), response_cost=0.0) is False
prisma.db.litellm_spendlogs.update_many.assert_not_called()
assert db_writer._batch_database_updates.await_count == 0
@pytest.mark.asyncio
@pytest.mark.parametrize(
("inserted", "existing", "charged"),
[
(1, None, True),
(0, SimpleNamespace(call_type="aretrieve_batch", status="success", spend=0.25), False),
],
ids=["first_retrieve_owns_the_row", "another_retrieve_already_charged"],
)
async def test_update_database_charges_a_batch_once_even_with_spend_logs_disabled(
inserted: int, existing: object, charged: bool
):
"""
disable_spend_logs drops the per-request logs, not the batch's charge, so the one row
that makes a batch chargeable exactly once is still written and still read back.
"""
db_writer = DBSpendUpdateWriter()
db_writer._batch_database_updates = AsyncMock()
prisma = _spend_logs_prisma(inserted, existing)
assert await _update_database_with(db_writer, prisma, _batch_cost_payload(), True) is charged
assert prisma.db.litellm_spendlogs.create_many.await_count == 1
assert db_writer._batch_database_updates.await_count == (1 if charged else 0)
@pytest.mark.asyncio
async def test_update_database_writes_no_ordinary_spend_row_with_spend_logs_disabled():
"""The batch carve-out above stays a carve-out: every other row still goes unwritten."""
db_writer = DBSpendUpdateWriter()
db_writer._batch_database_updates = AsyncMock()
prisma = _spend_logs_prisma(1, None)
payload = {**_batch_cost_payload(), "call_type": "acompletion"}
assert await _update_database_with(db_writer, prisma, payload, True) is True
prisma.db.litellm_spendlogs.create_many.assert_not_called()
assert prisma.spend_log_transactions == []
assert db_writer._batch_database_updates.await_count == 1
_BATCH_CLAIM_FIELDS = {"request_id", "call_type", "status", "spend", "startTime", "endTime"}
def _logged_batch_cost_payload() -> dict:
return {
**_batch_cost_payload(),
"api_key": "0e5b0e9e5f",
"model": "gpt-5.6-luna",
"user": "test-user",
"metadata": '{"batch_models": ["gpt-5.6-luna"]}',
"requester_ip_address": "127.0.0.1",
"proxy_server_request": '{"headers": {"user-agent": "litellm-batch-cost-check"}}',
}
@pytest.mark.asyncio
@pytest.mark.parametrize(
("disable_spend_logs", "logs_the_request"),
[(False, True), (True, False)],
ids=["spend_logs_on", "spend_logs_off"],
)
async def test_update_database_claims_a_batch_without_logging_the_request_that_polled_it(
disable_spend_logs: bool, logs_the_request: bool
):
"""
disable_spend_logs has to keep meaning that no request gets logged, and the batch's cost
row is the one row it cannot drop, so with logging off that row carries only what tells
the retrieves apart: no metadata, no requester IP, no key, model, or token counts.
"""
db_writer = DBSpendUpdateWriter()
db_writer._batch_database_updates = AsyncMock()
prisma = _spend_logs_prisma(1, None)
payload = _logged_batch_cost_payload()
assert await _update_database_with(db_writer, prisma, payload, disable_spend_logs) is True
claimed = prisma.db.litellm_spendlogs.create_many.await_args.kwargs["data"][0]
assert set(claimed) == (set(payload) if logs_the_request else _BATCH_CLAIM_FIELDS)
assert claimed["spend"] == 0.25
assert db_writer._batch_database_updates.await_count == 1
@pytest.mark.asyncio
async def test_update_database_queues_only_the_claim_for_a_batch_it_could_not_write_with_logs_disabled():
"""A refused claim is retried through the queue, so what it queues has to stay unlogged too."""
db_writer = DBSpendUpdateWriter()
db_writer._batch_database_updates = AsyncMock()
prisma = _spend_logs_prisma(0, None)
prisma.db.litellm_spendlogs.create_many = AsyncMock(side_effect=RuntimeError("db unreachable"))
assert await _update_database_with(db_writer, prisma, _logged_batch_cost_payload(), True) is True
assert [set(row) for row in prisma.spend_log_transactions] == [_BATCH_CLAIM_FIELDS]
assert db_writer._batch_database_updates.await_count == 1
@pytest.mark.asyncio
async def test_update_database_queues_a_batch_cost_row_it_could_not_claim():
"""An unreachable DB must not drop the batch's only spend row, nor its charge."""
db_writer = DBSpendUpdateWriter()
db_writer._batch_database_updates = AsyncMock()
prisma = _spend_logs_prisma(0, None)
prisma.db.litellm_spendlogs.create_many = AsyncMock(side_effect=RuntimeError("db unreachable"))
assert await _update_database_with(db_writer, prisma, _batch_cost_payload()) is True
assert [row["request_id"] for row in prisma.spend_log_transactions] == ["batch_abc_batch_cost"]
assert db_writer._batch_database_updates.await_count == 1
@pytest.mark.asyncio
@pytest.mark.parametrize(
"payload",
[{**_batch_cost_payload(), "call_type": "acompletion"}, {**_batch_cost_payload(), "status": "failure"}],
ids=["not_a_batch_retrieve", "failed_batch_retrieve"],
)
async def test_update_database_queues_every_other_spend_row_for_the_next_flush(payload: dict):
db_writer = DBSpendUpdateWriter()
db_writer._batch_database_updates = AsyncMock()
prisma = _spend_logs_prisma(1, None)
assert await _update_database_with(db_writer, prisma, payload) is True
prisma.db.litellm_spendlogs.create_many.assert_not_called()
assert prisma.spend_log_transactions == [payload]
assert db_writer._batch_database_updates.await_count == 1
@pytest.mark.asyncio
@pytest.mark.parametrize(
"injected_deployment, attributed",

View file

@ -1,4 +1,4 @@
import asyncio
from datetime import datetime
from unittest.mock import AsyncMock, MagicMock, patch
@ -70,9 +70,7 @@ async def test_async_post_call_failure_hook():
# Check that metadata was properly updated
assert "litellm_params" in call_args["kwargs"]
assert call_args["kwargs"]["litellm_params"]["proxy_server_request"] == {
"request_id": "test_request_id"
}
assert call_args["kwargs"]["litellm_params"]["proxy_server_request"] == {"request_id": "test_request_id"}
metadata = call_args["kwargs"]["litellm_params"]["metadata"]
assert metadata["user_api_key"] == "test_api_key"
assert metadata["status"] == "failure"
@ -336,9 +334,7 @@ async def test_should_continue_failure_tracking_when_budget_release_fails():
)
assert mock_invalidate_budget_reservation_counters.await_count == 1
assert (
mock_invalidate_budget_reservation_counters.await_args.kwargs[
"budget_reservation"
]
mock_invalidate_budget_reservation_counters.await_args.kwargs["budget_reservation"]
is user_api_key_dict.budget_reservation
)
assert user_api_key_dict.budget_reservation["finalized"] is True
@ -433,36 +429,21 @@ def test_get_budget_reservation_from_metadata_handles_dict_auth_object():
"entries": [{"counter_key": "spend:key:test_api_key"}],
}
assert _get_budget_reservation_from_metadata(metadata={"user_api_key_auth": dict(UserAPIKeyAuth())}) is None
assert (
_get_budget_reservation_from_metadata(
metadata={"user_api_key_auth": dict(UserAPIKeyAuth())}
)
is None
)
assert (
_get_budget_reservation_from_metadata(
metadata={
"user_api_key_auth": UserAPIKeyAuth(
budget_reservation=budget_reservation
)
}
metadata={"user_api_key_auth": UserAPIKeyAuth(budget_reservation=budget_reservation)}
)
== budget_reservation
)
assert (
_get_budget_reservation_from_metadata(
metadata={
"user_api_key_auth": dict(
UserAPIKeyAuth(budget_reservation=budget_reservation)
)
}
metadata={"user_api_key_auth": dict(UserAPIKeyAuth(budget_reservation=budget_reservation))}
)
== budget_reservation
)
assert (
_get_budget_reservation_from_metadata(
metadata={"user_api_key_budget_reservation": budget_reservation}
)
_get_budget_reservation_from_metadata(metadata={"user_api_key_budget_reservation": budget_reservation})
is budget_reservation
)
@ -470,9 +451,7 @@ def test_get_budget_reservation_from_metadata_handles_dict_auth_object():
@pytest.mark.asyncio
async def test_update_database_and_spend_counters_releases_reservation_when_db_update_fails():
proxy_logging_obj = MagicMock()
proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock(
side_effect=Exception("db unavailable")
)
proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock(side_effect=Exception("db unavailable"))
increment_spend_counters = AsyncMock()
budget_reservation = {"reserved_cost": 0.5, "entries": []}
@ -508,9 +487,7 @@ async def test_update_database_and_spend_counters_releases_reservation_when_db_u
async def test_update_database_and_spend_counters_preserves_db_exception_when_release_fails():
proxy_logging_obj = MagicMock()
db_exception = RuntimeError("db unavailable")
proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock(
side_effect=db_exception
)
proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock(side_effect=db_exception)
increment_spend_counters = AsyncMock()
budget_reservation = {"reserved_cost": 0.5, "entries": []}
@ -554,12 +531,8 @@ async def test_update_database_and_spend_counters_preserves_db_exception_when_re
budget_reservation=budget_reservation,
)
assert mock_log_exception.call_count == 2
mock_log_exception.assert_any_call(
"Failed to release budget reservation after database update failed"
)
mock_log_exception.assert_any_call(
"Failed to invalidate budget reservation counters after release failed"
)
mock_log_exception.assert_any_call("Failed to release budget reservation after database update failed")
mock_log_exception.assert_any_call("Failed to invalidate budget reservation counters after release failed")
increment_spend_counters.assert_not_awaited()
@ -778,6 +751,107 @@ async def test_track_cost_callback_defers_in_progress_background_interaction():
mock_proxy_logging.failed_tracking_alert.assert_not_called()
def _batch_retrieve_kwargs(call_type: str, reservation: dict | None = None) -> dict:
metadata = {
"user_api_key": "hashed_key",
"user_api_key_user_id": "user-1",
"user_api_key_team_id": "team-1",
**({"user_api_key_budget_reservation": reservation} if reservation is not None else {}),
}
return {
"call_type": call_type,
"model": "gpt-5.6-luna",
"litellm_call_id": "test-call-id",
"litellm_params": {"metadata": metadata},
"standard_logging_object": {"response_cost": 0.0, "request_tags": None},
"stream": False,
}
def _retrieved_batch(status: str, output_file_id: str | None):
from litellm.types.utils import LiteLLMBatch
return LiteLLMBatch(
id="batch_abc",
completion_window="24h",
created_at=1,
endpoint="/v1/chat/completions",
input_file_id="file-in",
object="batch",
status=status,
output_file_id=output_file_id,
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("call_type", "status", "output_file_id", "row_claimed", "spend_written", "charged"),
[
("aretrieve_batch", "in_progress", None, True, False, False),
("aretrieve_batch", "completed", None, True, False, False),
("aretrieve_batch", "completed", "file-out", False, True, False),
("aretrieve_batch", "completed", "file-out", True, True, True),
("aretrieve_batch", "failed", None, True, True, True),
("acreate_batch", "validating", None, True, True, True),
],
ids=[
"retrieve_before_final",
"retrieve_completed_without_output_yet",
"retrieve_after_another_retrieve_charged",
"retrieve_first_final",
"retrieve_failed_batch",
"create_before_final",
],
)
async def test_track_cost_callback_charges_a_batch_once_and_only_when_final( # test-quality-ok: whether the spend writer runs, whether the counters move, and whether the poll's reservation is handed back is the whole observable contract of the gate
call_type, status, output_file_id, row_claimed, spend_written, charged
):
"""
A poll before the batch is final used to pin its shared spend row at $0, and every
completed retrieve after the first charged the key again (LIT-7048). Only retrieves
are gated, since creating a batch is its own billable request, and a retrieve that
charges nothing hands its budget reservation back instead.
"""
logger = _ProxyDBLogger()
budget_reservation = None if charged else {"reserved_cost": 0.5, "entries": []}
kwargs = _batch_retrieve_kwargs(call_type, reservation=budget_reservation)
with (
patch( # test-quality-ok: increment_spend_counters is a proxy_server global the callback reads lazily, no seam
"litellm.proxy.proxy_server.increment_spend_counters", new_callable=AsyncMock
) as mock_increment_spend_counters,
patch( # test-quality-ok: update_cache is a proxy_server global the callback reads lazily, no seam
"litellm.proxy.proxy_server.update_cache", new_callable=AsyncMock
) as mock_update_cache,
patch( # test-quality-ok: callback imports proxy_logging_obj off proxy_server in its body, no seam
"litellm.proxy.proxy_server.proxy_logging_obj"
) as mock_proxy_logging,
patch( # test-quality-ok: the release is imported inside the callback's helper, no seam
"litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation", new_callable=AsyncMock
) as mock_release_budget_reservation,
):
mock_proxy_logging.failed_tracking_alert = AsyncMock()
mock_proxy_logging.db_spend_update_writer.update_database = AsyncMock(return_value=row_claimed)
mock_proxy_logging.slack_alerting_instance.customer_spend_alert = AsyncMock()
await logger._PROXY_track_cost_callback(
kwargs=kwargs,
completion_response=_retrieved_batch(status, output_file_id),
start_time=datetime.now(),
end_time=datetime.now(),
)
await asyncio.sleep(0)
mock_proxy_logging.failed_tracking_alert.assert_not_called()
assert mock_proxy_logging.db_spend_update_writer.update_database.await_count == (1 if spend_written else 0)
assert mock_increment_spend_counters.await_count == (1 if charged else 0)
assert mock_update_cache.await_count == (1 if charged else 0)
if charged:
mock_release_budget_reservation.assert_not_awaited()
else:
mock_release_budget_reservation.assert_awaited_once_with(budget_reservation=budget_reservation)
def _in_progress_interaction_kwargs(reservation: dict) -> dict:
return {
"call_type": "acreate_interaction",
@ -1101,10 +1175,7 @@ async def test_async_post_call_failure_hook_propagates_trace_id_from_logging_obj
# standard_logging_object should have been propagated from logging obj
assert call_kwargs.get("standard_logging_object") is not None
assert (
call_kwargs["standard_logging_object"]["trace_id"]
== "trace-id-from-logging-obj"
)
assert call_kwargs["standard_logging_object"]["trace_id"] == "trace-id-from-logging-obj"
# litellm_trace_id should also be propagated as a fallback
assert call_kwargs.get("litellm_trace_id") == "trace-id-from-logging-obj"
@ -1691,9 +1762,7 @@ async def test_async_post_call_failure_hook_records_recovered_partial_spend():
"metadata": {},
"proxy_server_request": {"request_id": "rid"},
"response_cost": 3.5e-05,
"combined_usage_object": Usage(
prompt_tokens=30, completion_tokens=1, total_tokens=31
),
"combined_usage_object": Usage(prompt_tokens=30, completion_tokens=1, total_tokens=31),
}
with patch(
@ -1772,15 +1841,10 @@ async def test_track_cost_callback_enriches_user_id_for_mcp_style_metadata():
assert mock_increment.call_args.kwargs["team_id"] == "team-123"
assert mock_increment.call_args.kwargs["org_id"] == "org-456"
update_kwargs = (
mock_proxy_logging.db_spend_update_writer.update_database.await_args.kwargs
)
update_kwargs = mock_proxy_logging.db_spend_update_writer.update_database.await_args.kwargs
assert update_kwargs["user_id"] == "mcp-user@example.com"
assert update_kwargs["team_id"] == "team-123"
assert (
kwargs["litellm_params"]["metadata"]["user_api_key_user_id"]
== "mcp-user@example.com"
)
assert kwargs["litellm_params"]["metadata"]["user_api_key_user_id"] == "mcp-user@example.com"
@pytest.mark.asyncio
@ -1875,9 +1939,7 @@ def test_should_track_cost_callback_pass_through_without_owner(call_type, expect
],
)
@pytest.mark.asyncio
async def test_track_cost_callback_logs_unauthenticated_pass_through_request(
call_type, expect_spend_log
):
async def test_track_cost_callback_logs_unauthenticated_pass_through_request(call_type, expect_spend_log):
"""Regression for LIT-3782: a pass-through request with auth=false reaches the
cost callback with no key/user/team/end-user. Before the fix the spend-log
write was skipped and the request never appeared in request/usage logs. It
@ -1923,9 +1985,7 @@ async def test_track_cost_callback_logs_unauthenticated_pass_through_request(
end_time=datetime.now(),
)
assert mock_proxy_logging.db_spend_update_writer.update_database.await_count == (
1 if expect_spend_log else 0
)
assert mock_proxy_logging.db_spend_update_writer.update_database.await_count == (1 if expect_spend_log else 0)
class _FakeDeploymentLookup:

View file

@ -389,14 +389,24 @@ class TestGpt6AstraAdvertisesItsDocumentedLevels:
"max",
)
@pytest.mark.parametrize("model", ["azure/gpt-6-astra", "azure/us/gpt-6-astra"])
def test_a_foundry_deployment_also_advertises_none(self, local_model_cost_map, model):
"""Microsoft Foundry serves the same model but its API accepts reasoning_effort none
(verified live: 200 with zero reasoning tokens, and it unlocks temperature), which
OpenAI's rejects, so an Azure deployment offers none on top of low through max."""
@pytest.mark.parametrize(
"model,custom_llm_provider",
[
("azure/gpt-6-astra", "azure"),
("azure/us/gpt-6-astra", "azure"),
("azure_ai/gpt-6-astra", "azure_ai"),
],
)
def test_an_azure_hosted_deployment_advertises_none_but_not_max(
self, local_model_cost_map, model, custom_llm_provider
):
"""Microsoft hosts the same model with a different level set than OpenAI does. Verified live
on both Azure routes: none returns 200 with zero reasoning tokens and unlocks temperature,
which OpenAI's API rejects, while max returns 400 unsupported_value naming none through
xhigh as the levels it does take."""
from litellm.utils import _get_model_info_helper
model_info = dict(_get_model_info_helper(model=model, custom_llm_provider="azure"))
model_info = dict(_get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider))
assert resolve_supported_reasoning_efforts(model_info, deployment_is_mapped=True) == (
"none",
@ -404,5 +414,4 @@ class TestGpt6AstraAdvertisesItsDocumentedLevels:
"medium",
"high",
"xhigh",
"max",
)