mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
merge: bring litellm_internal_staging into litellm_mcp_lifecycle_e2e
This commit is contained in:
commit
d85023e38c
18 changed files with 1003 additions and 104 deletions
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue