From 4c00a6e189d95f34a72036802937b38769f561b5 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 5 Sep 2026 16:35:32 -0700 Subject: [PATCH 01/15] fix(batches): account a batch's cost once, from the first retrieve that sees it final Every retrieve of a batch through the proxy shares one spend row, the batch id plus the batch cost suffix, and spend log inserts skip duplicates. A poll that landed while the batch was still validating or in progress wrote that row at $0 and no later retrieve could overwrite it, and every completed retrieve after the first added the cost to the key, team, and user counters again with no new row to show for it. The cost callback now writes nothing for a batch retrieve until the batch is final, releasing the poll's budget reservation instead, and once it is final it charges only when no spend row for that batch is queued for flush or already stored. Batch cost rows are flushed to the database right away so a second instance sees them, and the logger prices a batch only once it is final, which also covers a failed batch that never produced an output file. --- litellm/batches/batch_utils.py | 20 ++ litellm/litellm_core_utils/litellm_logging.py | 13 +- litellm/proxy/db/db_spend_update_writer.py | 3 +- .../proxy/hooks/proxy_track_cost_callback.py | 56 ++++- .../openai_files_endpoints/common_utils.py | 8 +- litellm/proxy/utils.py | 10 +- .../test_litellm/batches/test_batch_utils.py | 55 +++- .../test_litellm_logging.py | 82 ++++++ .../proxy/db/test_db_spend_update_writer.py | 10 +- .../hooks/test_proxy_track_cost_callback.py | 238 +++++++++++++----- .../prisma_and_spend/test_spend_functions.py | 15 ++ 11 files changed, 428 insertions(+), 82 deletions(-) diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index 97be5f77d79..eaac3bf0e9f 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -25,6 +25,26 @@ class BatchCostUsageResult: failed_requests: int +_TERMINAL_BATCH_STATUSES: Final = frozenset({"completed", "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 != "completed" 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"], diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index c31c4323157..09ddd1b9720 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -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 diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index e6880d521f1..ff48b00dc70 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -82,6 +82,7 @@ else: RESPONSES_SESSION_CALL_TYPES: Final = frozenset({CallTypes.responses.value, CallTypes.aresponses.value}) +IMMEDIATE_FLUSH_CALL_TYPES: Final = RESPONSES_SESSION_CALL_TYPES | frozenset({CallTypes.aretrieve_batch.value}) class _SpendBatch(Protocol): @@ -939,7 +940,7 @@ class DBSpendUpdateWriter: from litellm.proxy.utils import enqueue_spend_logs, request_spend_log_flush await enqueue_spend_logs(prisma_client, (payload,)) - if payload.get("call_type") in RESPONSES_SESSION_CALL_TYPES: + if payload.get("call_type") in IMMEDIATE_FLUSH_CALL_TYPES: request_spend_log_flush() else: verbose_proxy_logger.debug("prisma_client is None. Skipping writing spend logs to db.") diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 7254b05db2e..95e61fdd98c 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -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 ( @@ -33,17 +34,21 @@ from litellm.proxy.spend_tracking.spend_log_error_logger import ( from litellm.proxy.spend_tracking.spend_tracking_utils import ( _sanitize_error_information_for_spend_logs, get_request_model_access_groups, + get_spend_logs_id, ) from litellm.proxy.utils import ProxyUpdateSpend from litellm.types.utils import ( CallTypes, + LiteLLMBatch, StandardLoggingPayload, StandardLoggingPayloadErrorInformation, ) from litellm.utils import get_end_user_id_for_cost_tracking if TYPE_CHECKING: - from litellm.proxy.utils import ProxyLogging + from prisma.types import LiteLLM_SpendLogsWhereUniqueInput + + from litellm.proxy.utils import PrismaClient, ProxyLogging _UNATTRIBUTED_TRACKABLE_CALL_TYPES: Final[frozenset[str]] = frozenset( { @@ -224,6 +229,7 @@ class _ProxyDBLogger(CustomLogger): ): from litellm.proxy.proxy_server import ( increment_spend_counters, + prisma_client, proxy_logging_obj, update_cache, ) @@ -248,6 +254,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 + ): + batch_spend_log_id: Final = get_spend_logs_id( + CallTypes.aretrieve_batch.value, completion_response.model_dump(), kwargs + ) + if not await _batch_cost_is_trackable_now( + batch=completion_response, spend_log_id=batch_spend_log_id, prisma_client=prisma_client + ): + 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)) @@ -491,6 +509,42 @@ def _write_spend_metadata_to_kwargs(kwargs: dict, metadata: dict) -> None: bucket[key] = value +async def _batch_cost_is_trackable_now( + batch: LiteLLMBatch, spend_log_id: str | None, prisma_client: "PrismaClient | None" +) -> bool: + """A batch is billed exactly once, from the first retrieve that sees it final. + + Every retrieve of one batch shares a single spend row (its id plus the batch cost + suffix), so a poll that lands before the output exists would write that row at $0 + and pin it there, and every retrieve after the first would add the cost to the + key, team, and user counters again. + """ + if not batch_cost_is_final(batch): + verbose_proxy_logger.debug("Cost tracking deferred for batch %s still in status %s", batch.id, batch.status) + return False + if prisma_client is None or spend_log_id is None: + return True + if not await _spend_log_already_recorded(prisma_client=prisma_client, request_id=spend_log_id): + return True + verbose_proxy_logger.debug( + "Cost tracking skipped for batch %s: spend row %s already recorded", batch.id, spend_log_id + ) + return False + + +async def _spend_log_already_recorded(prisma_client: "PrismaClient", request_id: str) -> bool: + from litellm.proxy.utils import spend_log_is_queued + + if await spend_log_is_queued(prisma_client, request_id): + return True + spend_log_row: Final[LiteLLM_SpendLogsWhereUniqueInput] = {"request_id": request_id} + try: + return await prisma_client.db.litellm_spendlogs.find_unique(where=spend_log_row) is not None + except Exception as e: # noqa: BLE001 # prisma raises its own hierarchy; an unreadable DB must not drop the batch's only spend row + verbose_proxy_logger.warning("Could not check for an existing spend row %s, tracking anyway: %s", request_id, e) + return False + + def _is_unbilled_interaction_response(completion_response: object) -> bool: from litellm.interactions.background_cost_polling import missing_usage_is_expected from litellm.types.interactions import InteractionsAPIResponse diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index 15eeddbc489..b1f282a0978 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -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( diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index accf7b720fb..f64b51bc6c6 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -6251,7 +6251,9 @@ def request_spend_log_flush() -> None: The Responses API hands the client an id it can chain from straight away, and that lookup reads the DB, so the row cannot sit in this worker's queue for a poll interval. - Repeated requests coalesce into the monitor's next pass, so the batching holds. + A batch's cost row is what every other worker checks before charging the same batch + again, so it cannot wait either. Repeated requests coalesce into the monitor's next + pass, so the batching holds. """ PrismaClient.spend_log_flush_requested.set() @@ -6266,6 +6268,12 @@ async def _wait_for_spend_log_flush_request(interval: float) -> bool: return True +async def spend_log_is_queued(prisma_client: PrismaClient, request_id: str) -> bool: + """Whether a spend log with ``request_id`` is still waiting for the next flush.""" + async with prisma_client._spend_log_transactions_lock: + return any(row.get("request_id") == request_id for row in prisma_client.spend_log_transactions) + + async def dequeue_spend_logs(prisma_client: PrismaClient, limit: int) -> list[dict[str, object]]: """Take up to ``limit`` of the oldest queued spend logs off the queue. diff --git a/tests/test_litellm/batches/test_batch_utils.py b/tests/test_litellm/batches/test_batch_utils.py index c86c7c4df03..8d4f68164b4 100644 --- a/tests/test_litellm/batches/test_batch_utils.py +++ b/tests/test_litellm/batches/test_batch_utils.py @@ -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,55 @@ 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=status, + output_file_id=output_file_id, + request_counts=counts, + ) + + +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 + + def test_completed_with_output_is_final(self): + assert bu.batch_cost_is_final(_retrieved_batch("completed", 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 + + def test_completed_without_output_and_every_line_failed_is_final(self): + counts = BatchRequestCounts(total=2, completed=0, failed=2) + assert bu.batch_cost_is_final(_retrieved_batch("completed", 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 diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index 16a99713a06..4583429bd11 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -632,6 +632,88 @@ 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=status, + output_file_id=output_file_id, + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("status", "output_file_id"), + [("validating", None), ("in_progress", None), ("finalizing", None), ("completed", 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) + + with contextlib.suppress(Exception): + 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") + + with contextlib.suppress(Exception): + 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.""" diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index 11ef911de3e..e1b2d151c6d 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -2934,12 +2934,16 @@ async def test_commit_spend_updates_retries_deadlock_on_every_entity_path(monkey @pytest.mark.asyncio @pytest.mark.parametrize( "call_type, expects_flush", - [("aresponses", True), ("responses", True), ("acompletion", False)], + [("aresponses", True), ("responses", True), ("aretrieve_batch", 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. + Responses row cannot sit in this worker's queue until the monitor's next poll. A + batch's cost row is what another worker checks before charging the same batch again + (LIT-7048), so it cannot wait either. """ from litellm.proxy.utils import PrismaClient diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index 8043a1aca3f..b7037e8d621 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -1,4 +1,3 @@ - from datetime import datetime from unittest.mock import AsyncMock, MagicMock, patch @@ -7,6 +6,7 @@ import pytest from litellm.litellm_core_utils.internal_call_metadata import MODEL_ACCESS_GROUP_METADATA_KEY from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.proxy_track_cost_callback import ( + _batch_cost_is_trackable_now, _get_budget_reservation_from_metadata, _ProxyDBLogger, _should_track_cost_callback, @@ -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,169 @@ 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, + ) + + +def _prisma_client_with(queued_request_ids: tuple[str, ...], stored_row: object) -> MagicMock: + import asyncio + + prisma_client = MagicMock() + prisma_client._spend_log_transactions_lock = asyncio.Lock() + prisma_client.spend_log_transactions = [{"request_id": request_id} for request_id in queued_request_ids] + prisma_client.db.litellm_spendlogs.find_unique = AsyncMock(return_value=stored_row) + return prisma_client + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("status", "output_file_id", "spend_log_id", "prisma_client", "trackable"), + [ + ("in_progress", None, "batch_abc_batch_cost", None, False), + ("in_progress", None, "batch_abc_batch_cost", _prisma_client_with((), None), False), + ("completed", None, "batch_abc_batch_cost", _prisma_client_with((), None), False), + ("completed", "file-out", "batch_abc_batch_cost", None, True), + ("completed", "file-out", None, _prisma_client_with((), None), True), + ("completed", "file-out", "batch_abc_batch_cost", _prisma_client_with(("batch_abc_batch_cost",), None), False), + ( + "completed", + "file-out", + "batch_abc_batch_cost", + _prisma_client_with((), {"request_id": "batch_abc_batch_cost"}), + False, + ), + ("completed", "file-out", "batch_abc_batch_cost", _prisma_client_with((), None), True), + ("failed", None, "batch_abc_batch_cost", _prisma_client_with((), None), True), + ], + ids=[ + "in_progress_without_db", + "in_progress_never_consults_db", + "completed_without_output_yet", + "final_without_db", + "final_without_spend_log_id", + "final_row_queued_for_flush", + "final_row_already_stored", + "final_first_sighting", + "failed_first_sighting", + ], +) +async def test_batch_cost_is_trackable_now(status, output_file_id, spend_log_id, prisma_client, trackable): + """ + A batch is billed from the first retrieve that sees it final and never again: + a poll before that wrote the shared spend row at $0 and pinned it there, and + every completed retrieve after the first charged the key again (LIT-7048). + """ + assert ( + await _batch_cost_is_trackable_now( + batch=_retrieved_batch(status, output_file_id), spend_log_id=spend_log_id, prisma_client=prisma_client + ) + is trackable + ) + + +@pytest.mark.asyncio +async def test_batch_cost_is_trackable_now_when_the_spend_row_lookup_fails(): + """An unreadable spend log table must not drop the batch's only spend row.""" + prisma_client = _prisma_client_with((), None) + prisma_client.db.litellm_spendlogs.find_unique = AsyncMock(side_effect=RuntimeError("db unreachable")) + + assert ( + await _batch_cost_is_trackable_now( + batch=_retrieved_batch("completed", "file-out"), + spend_log_id="batch_abc_batch_cost", + prisma_client=prisma_client, + ) + is True + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("call_type", "status", "output_file_id", "stored_row", "charged"), + [ + ("aretrieve_batch", "in_progress", None, None, False), + ("aretrieve_batch", "completed", "file-out", {"request_id": "batch_abc_batch_cost"}, False), + ("aretrieve_batch", "completed", "file-out", None, True), + ("acreate_batch", "validating", None, None, True), + ], + ids=["retrieve_before_final", "retrieve_already_recorded", "retrieve_first_final", "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 and whether the poll's reservation is handed back is the whole observable contract of the gate + call_type, status, output_file_id, stored_row, charged +): + """ + Only retrieves are gated, since creating a batch is its own billable request. + A retrieve that writes 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: prisma_client is a proxy_server global the callback reads lazily, no seam + "litellm.proxy.proxy_server.prisma_client", _prisma_client_with((), stored_row) + ), + 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 + ), + 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 + ), + 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() + 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(), + ) + + mock_proxy_logging.failed_tracking_alert.assert_not_called() + if charged: + mock_proxy_logging.db_spend_update_writer.update_database.assert_awaited_once() + mock_release_budget_reservation.assert_not_awaited() + else: + mock_proxy_logging.db_spend_update_writer.update_database.assert_not_called() + 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 +1237,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 +1824,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 +1903,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.parametrize( @@ -1828,9 +1954,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 @@ -1876,9 +2000,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: diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py index a1eb88a7834..fc97d760226 100644 --- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py @@ -6,6 +6,7 @@ Symbols pinned here: - ``update_spend_logs_job`` - ``_monitor_spend_logs_queue`` - ``_raise_failed_update_spend_exception`` + - ``spend_log_is_queued`` """ from __future__ import annotations @@ -22,6 +23,7 @@ from litellm.proxy.utils import ( _monitor_spend_logs_queue, _raise_failed_update_spend_exception, drain_spend_logs_queue, + spend_log_is_queued, update_daily_tag_spend, update_spend, update_spend_logs_job, @@ -629,3 +631,16 @@ def test_raise_failed_update_spend_exception_raises_original_error() -> None: with pytest.raises(ValueError, match="specific"): asyncio.run(_runner()) + + +@pytest.mark.asyncio +async def test_spend_log_is_queued_matches_only_rows_awaiting_flush( + mock_prisma_client: Any, make_spend_log_row: Any +) -> None: + mock_prisma_client.spend_log_transactions = [make_spend_log_row(request_id="batch_abc_batch_cost")] + + assert await spend_log_is_queued(mock_prisma_client, "batch_abc_batch_cost") is True + assert await spend_log_is_queued(mock_prisma_client, "batch_abc") is False + + mock_prisma_client.spend_log_transactions = [] + assert await spend_log_is_queued(mock_prisma_client, "batch_abc_batch_cost") is False From 635bb3a2096fbb4f4c8899574807c049bb8e4825 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 5 Sep 2026 17:08:42 -0700 Subject: [PATCH 02/15] feat(cost-map): add azure_ai/gpt-6-astra Foundry pricing A gpt-6-astra deployment on a Foundry project reached through the azure_ai route had no cost map entry of its own, so it resolved to the OpenAI gpt-6-astra card: missing from the azure_ai/* wildcard listing, flex and priority prices and /v1/batch it does not sell, and no none reasoning effort. Add azure_ai/gpt-6-astra mirroring the azure/gpt-6-astra Standard Global sheet the way azure_ai/gpt-5.5 mirrors azure/gpt-5.5, and extend the cost, reasoning-effort, and wildcard listing tests to the Foundry route. --- ...odel_prices_and_context_window_backup.json | 43 +++++++++++++++++++ model_prices_and_context_window.json | 43 +++++++++++++++++++ .../llm_cost_calc/test_llm_cost_calc_utils.py | 15 +++++-- .../proxy/auth/test_model_checks.py | 19 ++++++++ .../test_reasoning_effort_capability.py | 16 +++++-- 5 files changed, 129 insertions(+), 7 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 4273ec54472..ac7407c2608 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -3485,6 +3485,49 @@ "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, + "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": true, + "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_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, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 4273ec54472..ac7407c2608 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -3485,6 +3485,49 @@ "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, + "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": true, + "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_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, diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py index df680b7cb0e..73f1a19d85c 100644 --- a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py +++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py @@ -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 diff --git a/tests/test_litellm/proxy/auth/test_model_checks.py b/tests/test_litellm/proxy/auth/test_model_checks.py index d58683fd1e5..eb48f70d5da 100644 --- a/tests/test_litellm/proxy/auth/test_model_checks.py +++ b/tests/test_litellm/proxy/auth/test_model_checks.py @@ -857,6 +857,25 @@ 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): + """A Foundry (azure_ai) deployment of gpt-6-astra only shows up under an azure_ai/* wildcard + when the cost map carries its own azure_ai/ entry; the azure/ entry from the OpenAI-on-Azure + price sheet never reaches the Foundry provider list (LIT-7081).""" + 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] + try: + litellm.add_known_models(model_cost_map={foundry_key: local_entry}) + assert foundry_key in get_known_models_from_wildcard("azure_ai/*") + finally: + 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 diff --git a/tests/test_litellm/router_utils/test_reasoning_effort_capability.py b/tests/test_litellm/router_utils/test_reasoning_effort_capability.py index f181370455d..b7499d1c975 100644 --- a/tests/test_litellm/router_utils/test_reasoning_effort_capability.py +++ b/tests/test_litellm/router_utils/test_reasoning_effort_capability.py @@ -389,14 +389,22 @@ 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): + @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_a_foundry_deployment_also_advertises_none(self, local_model_cost_map, model, custom_llm_provider): """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.""" + OpenAI's rejects, so an Azure deployment offers none on top of low through max, whether + it is reached through the azure route or the azure_ai (Foundry) route.""" 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", From b067e836f8df15bb3dbefe345bf503b7d3811b9a Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 5 Sep 2026 18:58:11 -0700 Subject: [PATCH 03/15] fix(batches): claim the batch cost spend row in the database before charging The cost callback used to look for an existing `_batch_cost` row before charging a completed batch, which left a window where concurrent retrieves on any instance all charged the key, and it would honor a row any request had written under that id. The spend update writer now inserts the batch cost row itself with `create_many(skip_duplicates=True)` and only the retrieve whose insert lands charges the key, team, and user. An existing row only takes the charge when it is a successful `aretrieve_batch` row, so a client-chosen `x-litellm-call-id` on another endpoint cannot suppress billing. Batch cost rows no longer get their own immediate flush path `batch_cost_is_final` now treats the proxy's normalized `complete` status like `completed`, which the enterprise batch cost poller relies on when it decides whether a completed batch is safe to retire. Tests build that status with `model_copy` since the OpenAI `Batch` model rejects it The `test-quality-ok` markers sit on the `patch(` lines the gate keys on, and the logging tests no longer wrap the priced retrieve in `contextlib.suppress` --- litellm/batches/batch_utils.py | 5 +- litellm/proxy/db/db_spend_update_writer.py | 73 ++++++++-- .../proxy/hooks/proxy_track_cost_callback.py | 68 +++------- litellm/proxy/utils.py | 10 +- .../test_litellm/batches/test_batch_utils.py | 14 +- .../test_litellm_logging.py | 12 +- .../proxy/db/test_db_spend_update_writer.py | 128 +++++++++++++++++- .../hooks/test_proxy_track_cost_callback.py | 116 ++++------------ .../prisma_and_spend/test_spend_functions.py | 15 -- 9 files changed, 249 insertions(+), 192 deletions(-) diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index eaac3bf0e9f..959c7498479 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -25,7 +25,8 @@ class BatchCostUsageResult: failed_requests: int -_TERMINAL_BATCH_STATUSES: Final = frozenset({"completed", "failed", "cancelled", "expired"}) +_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: @@ -39,7 +40,7 @@ def batch_cost_is_final(batch: Batch) -> bool: """ if batch.status not in _TERMINAL_BATCH_STATUSES: return False - if batch.status != "completed" or batch.output_file_id is not None: + 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 diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index ff48b00dc70..3fad351224b 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -82,7 +82,10 @@ else: RESPONSES_SESSION_CALL_TYPES: Final = frozenset({CallTypes.responses.value, CallTypes.aresponses.value}) -IMMEDIATE_FLUSH_CALL_TYPES: Final = RESPONSES_SESSION_CALL_TYPES | frozenset({CallTypes.aretrieve_batch.value}) + + +def _is_batch_cost_row(payload: SpendLogsPayload) -> bool: + return payload.get("call_type") == CallTypes.aretrieve_batch.value and payload.get("status") == "success" class _SpendBatch(Protocol): @@ -216,7 +219,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, @@ -233,7 +241,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: @@ -264,10 +272,8 @@ class DBSpendUpdateWriter: payload["team_id"] = team_id if disable_spend_logs is False: - await self._insert_spend_log_to_db( - payload=payload, - prisma_client=prisma_client, - ) + if not await self._record_spend_log(payload=payload, prisma_client=prisma_client): + return False await self._enqueue_tool_usage_transaction( payload=payload, completion_response=completion_response, @@ -307,6 +313,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 " @@ -319,7 +326,55 @@ class DBSpendUpdateWriter: org_id, end_user_id, ) - return + return True + + async def _record_spend_log(self, payload: SpendLogsPayload, prisma_client: "PrismaClient | None") -> bool: + if prisma_client is None or not _is_batch_cost_row(payload): + await self._insert_spend_log_to_db(payload=payload, prisma_client=prisma_client) + return True + return await self._claim_batch_cost_spend_log(payload=payload, prisma_client=prisma_client) + + async def _claim_batch_cost_spend_log(self, payload: SpendLogsPayload, prisma_client: "PrismaClient") -> 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 a successful retrieve wrote counts: a failed retrieve, or any request whose + client picked the batch id as its call id, cannot take the charge away. + """ + from litellm.repositories.table_repositories import SpendLogsRepository + + request_id: Final = payload["request_id"] + spend_logs: Final = SpendLogsRepository(prisma_client).table + try: + claimed: Final = await spend_logs.create_many( + data=[prisma_client.jsonify_object(payload)], # 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=payload, prisma_client=prisma_client) + return True + if ( + existing is not None + and existing.call_type == CallTypes.aretrieve_batch.value + and existing.status == "success" + ): + verbose_proxy_logger.debug("Cost tracking skipped: spend row %s already charged this batch", request_id) + return False + 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 async def _enqueue_tool_usage_transaction( self, @@ -940,7 +995,7 @@ class DBSpendUpdateWriter: from litellm.proxy.utils import enqueue_spend_logs, request_spend_log_flush await enqueue_spend_logs(prisma_client, (payload,)) - if payload.get("call_type") in IMMEDIATE_FLUSH_CALL_TYPES: + if payload.get("call_type") in RESPONSES_SESSION_CALL_TYPES: request_spend_log_flush() else: verbose_proxy_logger.debug("prisma_client is None. Skipping writing spend logs to db.") diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 95e61fdd98c..f0c889a8cb4 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -34,7 +34,6 @@ from litellm.proxy.spend_tracking.spend_log_error_logger import ( from litellm.proxy.spend_tracking.spend_tracking_utils import ( _sanitize_error_information_for_spend_logs, get_request_model_access_groups, - get_spend_logs_id, ) from litellm.proxy.utils import ProxyUpdateSpend from litellm.types.utils import ( @@ -46,9 +45,7 @@ from litellm.types.utils import ( from litellm.utils import get_end_user_id_for_cost_tracking if TYPE_CHECKING: - from prisma.types import LiteLLM_SpendLogsWhereUniqueInput - - from litellm.proxy.utils import PrismaClient, ProxyLogging + from litellm.proxy.utils import ProxyLogging _UNATTRIBUTED_TRACKABLE_CALL_TYPES: Final[frozenset[str]] = frozenset( { @@ -229,7 +226,6 @@ class _ProxyDBLogger(CustomLogger): ): from litellm.proxy.proxy_server import ( increment_spend_counters, - prisma_client, proxy_logging_obj, update_cache, ) @@ -257,15 +253,15 @@ class _ProxyDBLogger(CustomLogger): if ( isinstance(completion_response, LiteLLMBatch) and kwargs.get("call_type") == CallTypes.aretrieve_batch.value + and not batch_cost_is_final(completion_response) ): - batch_spend_log_id: Final = get_spend_logs_id( - CallTypes.aretrieve_batch.value, completion_response.model_dump(), kwargs + verbose_proxy_logger.debug( + "Cost tracking deferred for batch %s still in status %s", + completion_response.id, + completion_response.status, ) - if not await _batch_cost_is_trackable_now( - batch=completion_response, spend_log_id=batch_spend_log_id, prisma_client=prisma_client - ): - await _release_budget_reservation(budget_reservation=budget_reservation) - return + 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)) @@ -307,7 +303,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, @@ -324,6 +320,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.) @@ -509,42 +507,6 @@ def _write_spend_metadata_to_kwargs(kwargs: dict, metadata: dict) -> None: bucket[key] = value -async def _batch_cost_is_trackable_now( - batch: LiteLLMBatch, spend_log_id: str | None, prisma_client: "PrismaClient | None" -) -> bool: - """A batch is billed exactly once, from the first retrieve that sees it final. - - Every retrieve of one batch shares a single spend row (its id plus the batch cost - suffix), so a poll that lands before the output exists would write that row at $0 - and pin it there, and every retrieve after the first would add the cost to the - key, team, and user counters again. - """ - if not batch_cost_is_final(batch): - verbose_proxy_logger.debug("Cost tracking deferred for batch %s still in status %s", batch.id, batch.status) - return False - if prisma_client is None or spend_log_id is None: - return True - if not await _spend_log_already_recorded(prisma_client=prisma_client, request_id=spend_log_id): - return True - verbose_proxy_logger.debug( - "Cost tracking skipped for batch %s: spend row %s already recorded", batch.id, spend_log_id - ) - return False - - -async def _spend_log_already_recorded(prisma_client: "PrismaClient", request_id: str) -> bool: - from litellm.proxy.utils import spend_log_is_queued - - if await spend_log_is_queued(prisma_client, request_id): - return True - spend_log_row: Final[LiteLLM_SpendLogsWhereUniqueInput] = {"request_id": request_id} - try: - return await prisma_client.db.litellm_spendlogs.find_unique(where=spend_log_row) is not None - except Exception as e: # noqa: BLE001 # prisma raises its own hierarchy; an unreadable DB must not drop the batch's only spend row - verbose_proxy_logger.warning("Could not check for an existing spend row %s, tracking anyway: %s", request_id, e) - return False - - def _is_unbilled_interaction_response(completion_response: object) -> bool: from litellm.interactions.background_cost_polling import missing_usage_is_expected from litellm.types.interactions import InteractionsAPIResponse @@ -636,9 +598,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, @@ -663,6 +625,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( @@ -688,6 +653,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: diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index f64b51bc6c6..accf7b720fb 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -6251,9 +6251,7 @@ def request_spend_log_flush() -> None: The Responses API hands the client an id it can chain from straight away, and that lookup reads the DB, so the row cannot sit in this worker's queue for a poll interval. - A batch's cost row is what every other worker checks before charging the same batch - again, so it cannot wait either. Repeated requests coalesce into the monitor's next - pass, so the batching holds. + Repeated requests coalesce into the monitor's next pass, so the batching holds. """ PrismaClient.spend_log_flush_requested.set() @@ -6268,12 +6266,6 @@ async def _wait_for_spend_log_flush_request(interval: float) -> bool: return True -async def spend_log_is_queued(prisma_client: PrismaClient, request_id: str) -> bool: - """Whether a spend log with ``request_id`` is still waiting for the next flush.""" - async with prisma_client._spend_log_transactions_lock: - return any(row.get("request_id") == request_id for row in prisma_client.spend_log_transactions) - - async def dequeue_spend_logs(prisma_client: PrismaClient, limit: int) -> list[dict[str, object]]: """Take up to ``limit`` of the oldest queued spend logs off the queue. diff --git a/tests/test_litellm/batches/test_batch_utils.py b/tests/test_litellm/batches/test_batch_utils.py index 8d4f68164b4..976a96f2db1 100644 --- a/tests/test_litellm/batches/test_batch_utils.py +++ b/tests/test_litellm/batches/test_batch_utils.py @@ -1735,10 +1735,10 @@ def _retrieved_batch( endpoint="/v1/chat/completions", input_file_id="file-in", object="batch", - status=status, + status="validating", output_file_id=output_file_id, request_counts=counts, - ) + ).model_copy(update={"status": status}) class TestBatchCostIsFinal: @@ -1750,8 +1750,9 @@ class TestBatchCostIsFinal: def test_in_flight_batch_is_not_final(self, status): assert bu.batch_cost_is_final(_retrieved_batch(status)) is False - def test_completed_with_output_is_final(self): - assert bu.batch_cost_is_final(_retrieved_batch("completed", output_file_id="file-out")) is True + @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 @@ -1764,9 +1765,10 @@ class TestBatchCostIsFinal: counts = BatchRequestCounts(total=2, completed=2, failed=0) assert bu.batch_cost_is_final(_retrieved_batch("completed", counts=counts)) is False - def test_completed_without_output_and_every_line_failed_is_final(self): + @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("completed", counts=counts)) is True + 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): diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index 4583429bd11..174efaa4679 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -665,14 +665,14 @@ class TestRetrieveBatchPricesOnlyFinalBatches: endpoint="/v1/chat/completions", input_file_id="file-in", object="batch", - status=status, + 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)], + [("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 @@ -681,8 +681,7 @@ class TestRetrieveBatchPricesOnlyFinalBatches: monkeypatch.setattr(logging_module, "_handle_completed_batch", handle_completed_batch) batch = self._batch(status, output_file_id) - with contextlib.suppress(Exception): - await self._logging_obj()._async_success_handler_body(result=batch, start_time=None, end_time=None) + 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 @@ -705,8 +704,7 @@ class TestRetrieveBatchPricesOnlyFinalBatches: monkeypatch.setattr(logging_module, "_handle_completed_batch", handle_completed_batch) batch = self._batch("completed", "file-out") - with contextlib.suppress(Exception): - await self._logging_obj()._async_success_handler_body(result=batch, start_time=None, end_time=None) + 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 diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index e1b2d151c6d..500a0e7bb06 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -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 @@ -2934,16 +2935,14 @@ async def test_commit_spend_updates_retries_deadlock_on_every_entity_path(monkey @pytest.mark.asyncio @pytest.mark.parametrize( "call_type, expects_flush", - [("aresponses", True), ("responses", True), ("aretrieve_batch", True), ("acompletion", False)], + [("aresponses", True), ("responses", True), ("acompletion", False)], ) 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. A - batch's cost row is what another worker checks before charging the same batch again - (LIT-7048), so it cannot wait either. + Responses row cannot sit in this worker's queue until the monitor's next poll. """ from litellm.proxy.utils import PrismaClient @@ -2961,6 +2960,127 @@ async def test_insert_spend_log_asks_for_an_immediate_flush_on_rows_other_worker 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) -> 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) + return prisma + + +async def _update_database_with(db_writer: DBSpendUpdateWriter, prisma: MagicMock, payload: dict) -> 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", False + ), + 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=0.25, + ) + 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"), False), + (0, SimpleNamespace(call_type="aretrieve_batch", status="failure"), True), + (0, SimpleNamespace(call_type="aembedding", status="success"), True), + (0, None, True), + ], + ids=[ + "first_retrieve_owns_the_row", + "another_retrieve_already_charged", + "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 + written by anything but a successful retrieve, say a request whose client picked the + batch id as its call id, must not be able to take the charge away. + """ + 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 +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", diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index b7037e8d621..2965f8b4006 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -1,3 +1,4 @@ +import asyncio from datetime import datetime from unittest.mock import AsyncMock, MagicMock, patch @@ -6,7 +7,6 @@ import pytest from litellm.litellm_core_utils.internal_call_metadata import MODEL_ACCESS_GROUP_METADATA_KEY from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.proxy_track_cost_callback import ( - _batch_cost_is_trackable_now, _get_budget_reservation_from_metadata, _ProxyDBLogger, _should_track_cost_callback, @@ -783,110 +783,46 @@ def _retrieved_batch(status: str, output_file_id: str | None): ) -def _prisma_client_with(queued_request_ids: tuple[str, ...], stored_row: object) -> MagicMock: - import asyncio - - prisma_client = MagicMock() - prisma_client._spend_log_transactions_lock = asyncio.Lock() - prisma_client.spend_log_transactions = [{"request_id": request_id} for request_id in queued_request_ids] - prisma_client.db.litellm_spendlogs.find_unique = AsyncMock(return_value=stored_row) - return prisma_client - - @pytest.mark.asyncio @pytest.mark.parametrize( - ("status", "output_file_id", "spend_log_id", "prisma_client", "trackable"), + ("call_type", "status", "output_file_id", "row_claimed", "spend_written", "charged"), [ - ("in_progress", None, "batch_abc_batch_cost", None, False), - ("in_progress", None, "batch_abc_batch_cost", _prisma_client_with((), None), False), - ("completed", None, "batch_abc_batch_cost", _prisma_client_with((), None), False), - ("completed", "file-out", "batch_abc_batch_cost", None, True), - ("completed", "file-out", None, _prisma_client_with((), None), True), - ("completed", "file-out", "batch_abc_batch_cost", _prisma_client_with(("batch_abc_batch_cost",), None), False), - ( - "completed", - "file-out", - "batch_abc_batch_cost", - _prisma_client_with((), {"request_id": "batch_abc_batch_cost"}), - False, - ), - ("completed", "file-out", "batch_abc_batch_cost", _prisma_client_with((), None), True), - ("failed", None, "batch_abc_batch_cost", _prisma_client_with((), None), True), + ("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=[ - "in_progress_without_db", - "in_progress_never_consults_db", - "completed_without_output_yet", - "final_without_db", - "final_without_spend_log_id", - "final_row_queued_for_flush", - "final_row_already_stored", - "final_first_sighting", - "failed_first_sighting", + "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_batch_cost_is_trackable_now(status, output_file_id, spend_log_id, prisma_client, trackable): - """ - A batch is billed from the first retrieve that sees it final and never again: - a poll before that wrote the shared spend row at $0 and pinned it there, and - every completed retrieve after the first charged the key again (LIT-7048). - """ - assert ( - await _batch_cost_is_trackable_now( - batch=_retrieved_batch(status, output_file_id), spend_log_id=spend_log_id, prisma_client=prisma_client - ) - is trackable - ) - - -@pytest.mark.asyncio -async def test_batch_cost_is_trackable_now_when_the_spend_row_lookup_fails(): - """An unreadable spend log table must not drop the batch's only spend row.""" - prisma_client = _prisma_client_with((), None) - prisma_client.db.litellm_spendlogs.find_unique = AsyncMock(side_effect=RuntimeError("db unreachable")) - - assert ( - await _batch_cost_is_trackable_now( - batch=_retrieved_batch("completed", "file-out"), - spend_log_id="batch_abc_batch_cost", - prisma_client=prisma_client, - ) - is True - ) - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - ("call_type", "status", "output_file_id", "stored_row", "charged"), - [ - ("aretrieve_batch", "in_progress", None, None, False), - ("aretrieve_batch", "completed", "file-out", {"request_id": "batch_abc_batch_cost"}, False), - ("aretrieve_batch", "completed", "file-out", None, True), - ("acreate_batch", "validating", None, None, True), - ], - ids=["retrieve_before_final", "retrieve_already_recorded", "retrieve_first_final", "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 and whether the poll's reservation is handed back is the whole observable contract of the gate - call_type, status, output_file_id, stored_row, charged +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 ): """ - Only retrieves are gated, since creating a batch is its own billable request. - A retrieve that writes nothing hands its budget reservation back instead. + 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: prisma_client is a proxy_server global the callback reads lazily, no seam - "litellm.proxy.proxy_server.prisma_client", _prisma_client_with((), stored_row) - ), 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, @@ -895,7 +831,7 @@ async def test_track_cost_callback_charges_a_batch_once_and_only_when_final( # ) as mock_release_budget_reservation, ): mock_proxy_logging.failed_tracking_alert = AsyncMock() - mock_proxy_logging.db_spend_update_writer.update_database = 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( @@ -904,13 +840,15 @@ async def test_track_cost_callback_charges_a_batch_once_and_only_when_final( # 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_proxy_logging.db_spend_update_writer.update_database.assert_awaited_once() mock_release_budget_reservation.assert_not_awaited() else: - mock_proxy_logging.db_spend_update_writer.update_database.assert_not_called() mock_release_budget_reservation.assert_awaited_once_with(budget_reservation=budget_reservation) diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py index fc97d760226..a1eb88a7834 100644 --- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py @@ -6,7 +6,6 @@ Symbols pinned here: - ``update_spend_logs_job`` - ``_monitor_spend_logs_queue`` - ``_raise_failed_update_spend_exception`` - - ``spend_log_is_queued`` """ from __future__ import annotations @@ -23,7 +22,6 @@ from litellm.proxy.utils import ( _monitor_spend_logs_queue, _raise_failed_update_spend_exception, drain_spend_logs_queue, - spend_log_is_queued, update_daily_tag_spend, update_spend, update_spend_logs_job, @@ -631,16 +629,3 @@ def test_raise_failed_update_spend_exception_raises_original_error() -> None: with pytest.raises(ValueError, match="specific"): asyncio.run(_runner()) - - -@pytest.mark.asyncio -async def test_spend_log_is_queued_matches_only_rows_awaiting_flush( - mock_prisma_client: Any, make_spend_log_row: Any -) -> None: - mock_prisma_client.spend_log_transactions = [make_spend_log_row(request_id="batch_abc_batch_cost")] - - assert await spend_log_is_queued(mock_prisma_client, "batch_abc_batch_cost") is True - assert await spend_log_is_queued(mock_prisma_client, "batch_abc") is False - - mock_prisma_client.spend_log_transactions = [] - assert await spend_log_is_queued(mock_prisma_client, "batch_abc_batch_cost") is False From 15372967c6cd5085d5d3d9ebb30c5a158f3f8170 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 5 Sep 2026 19:06:38 -0700 Subject: [PATCH 04/15] fix(azure_ai): read the azure_ai card for gpt-5 series reasoning effort gates Foundry deployments of gpt-6-astra reached through azure_ai used the bare OpenAI card for the reasoning_effort none gates, so temperature and top_p were refused while the azure_ai card says none is supported. AzureAIStudioConfig now dispatches gpt-5 series params through AzureAIGPT5Config, which looks capabilities up under the azure_ai/ prefix the way the azure route does Also carries the search_context_cost_per_query block azure/gpt-6-astra has, adds a flex service tier cost test that fails at the merge base, and keeps the wildcard test from stripping azure_ai/gpt-6-astra out of the provider set --- litellm/llms/azure_ai/chat/transformation.py | 37 ++++++++++++++++++- ...odel_prices_and_context_window_backup.json | 5 +++ model_prices_and_context_window.json | 5 +++ .../llm_cost_calc/test_llm_cost_calc_utils.py | 15 ++++++++ .../chat/test_azure_ai_transformation.py | 22 +++++++++++ .../proxy/auth/test_model_checks.py | 6 ++- 6 files changed, 87 insertions(+), 3 deletions(-) diff --git a/litellm/llms/azure_ai/chat/transformation.py b/litellm/llms/azure_ai/chat/transformation.py index f2d405e9a17..05abd5882c6 100644 --- a/litellm/llms/azure_ai/chat/transformation.py +++ b/litellm/llms/azure_ai/chat/transformation.py @@ -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,25 @@ NON_OPENAI_SPEC_MESSAGE_FIELDS: Final = ( ) +class AzureAIGPT5Config(OpenAIGPT5Config): + @classmethod + def _model_map_lookup_name(cls, model: str) -> str: + return model if model.startswith("azure_ai/") else f"azure_ai/{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 +75,27 @@ class AzureAIStudioConfig(OpenAIConfig): return supported_params + def map_openai_params( + self, + non_default_params: dict, # mutable-ok: OpenAIConfig.map_openai_params signature + optional_params: dict, # mutable-ok: OpenAIConfig.map_openai_params signature + model: str, + drop_params: bool, + ) -> dict: # 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. diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index ac7407c2608..5b4652d38c7 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -3499,6 +3499,11 @@ "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", diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index ac7407c2608..5b4652d38c7 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -3499,6 +3499,11 @@ "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", diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py index 73f1a19d85c..caf97ba791d 100644 --- a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py +++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py @@ -2060,6 +2060,21 @@ 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): + """Foundry sells gpt-6-astra on Standard Global only, so a flex service_tier bills the standard rate. + The bare OpenAI card the azure_ai route fell back to before this entry existed carries flex prices + at half rate (LIT-7081).""" + 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", [ diff --git a/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py b/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py index 33fbb4e8fc7..25eca3b37ad 100644 --- a/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py +++ b/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py @@ -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,26 @@ 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): + """A Foundry deployment reached through azure_ai reads the azure_ai/ card, where gpt-6-astra supports + reasoning_effort none, so temperature and top_p ride along; the bare OpenAI card says none is + unsupported and the route used to refuse temperature and drop top_p (LIT-7081).""" + 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_azure_ai_grok_stop_parameter_handling(): """ Test that Grok models properly handle stop parameter filtering in Azure AI Studio. diff --git a/tests/test_litellm/proxy/auth/test_model_checks.py b/tests/test_litellm/proxy/auth/test_model_checks.py index eb48f70d5da..56dbcca61f3 100644 --- a/tests/test_litellm/proxy/auth/test_model_checks.py +++ b/tests/test_litellm/proxy/auth/test_model_checks.py @@ -868,12 +868,14 @@ def test_azure_ai_wildcard_lists_the_foundry_gpt_6_astra_entry(monkeypatch): 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: - litellm.azure_ai_models.discard(foundry_key) - litellm.add_known_models(model_cost_map={}) + 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 a17fcecf7092d0333afed4168d1954a1e9675d18 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 5 Sep 2026 19:21:34 -0700 Subject: [PATCH 05/15] refactor(azure_ai): type the Foundry param mapping override and drop test docstrings The AzureAIStudioConfig.map_openai_params override now carries dict[str, object] annotations instead of bare dict, and the docstrings added to the new tests go away since the test names already say what they cover. No behavior change --- litellm/llms/azure_ai/chat/transformation.py | 6 +++--- .../llm_cost_calc/test_llm_cost_calc_utils.py | 3 --- .../llms/azure_ai/chat/test_azure_ai_transformation.py | 3 --- tests/test_litellm/proxy/auth/test_model_checks.py | 3 --- .../router_utils/test_reasoning_effort_capability.py | 3 +-- 5 files changed, 4 insertions(+), 14 deletions(-) diff --git a/litellm/llms/azure_ai/chat/transformation.py b/litellm/llms/azure_ai/chat/transformation.py index 05abd5882c6..7c9a26c3f07 100644 --- a/litellm/llms/azure_ai/chat/transformation.py +++ b/litellm/llms/azure_ai/chat/transformation.py @@ -77,11 +77,11 @@ class AzureAIStudioConfig(OpenAIConfig): def map_openai_params( self, - non_default_params: dict, # mutable-ok: OpenAIConfig.map_openai_params signature - optional_params: dict, # mutable-ok: OpenAIConfig.map_openai_params signature + 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: # mutable-ok: OpenAIConfig.map_openai_params signature + ) -> 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, diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py index caf97ba791d..40abb5bfca3 100644 --- a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py +++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py @@ -2061,9 +2061,6 @@ def test_generic_cost_per_token_azure_gpt_6_astra_foundry_price_sheet( def test_generic_cost_per_token_azure_ai_gpt_6_astra_flex_bills_the_standard_rate(_local_model_cost_map): - """Foundry sells gpt-6-astra on Standard Global only, so a flex service_tier bills the standard rate. - The bare OpenAI card the azure_ai route fell back to before this entry existed carries flex prices - at half rate (LIT-7081).""" 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") diff --git a/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py b/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py index 25eca3b37ad..5ff0b729449 100644 --- a/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py +++ b/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py @@ -147,9 +147,6 @@ def _local_model_cost_map(monkeypatch: pytest.MonkeyPatch) -> None: def test_foundry_gpt_6_astra_keeps_sampling_params_when_reasoning_effort_is_none(_local_model_cost_map): - """A Foundry deployment reached through azure_ai reads the azure_ai/ card, where gpt-6-astra supports - reasoning_effort none, so temperature and top_p ride along; the bare OpenAI card says none is - unsupported and the route used to refuse temperature and drop top_p (LIT-7081).""" optional_params = AzureAIStudioConfig().map_openai_params( non_default_params={"reasoning_effort": "none", "temperature": 0.2, "top_p": 0.9}, optional_params={}, diff --git a/tests/test_litellm/proxy/auth/test_model_checks.py b/tests/test_litellm/proxy/auth/test_model_checks.py index 56dbcca61f3..36bfc4c5dd3 100644 --- a/tests/test_litellm/proxy/auth/test_model_checks.py +++ b/tests/test_litellm/proxy/auth/test_model_checks.py @@ -859,9 +859,6 @@ def test_add_known_models_refreshes_models_by_provider_for_wildcard_expansion(): def test_azure_ai_wildcard_lists_the_foundry_gpt_6_astra_entry(monkeypatch): - """A Foundry (azure_ai) deployment of gpt-6-astra only shows up under an azure_ai/* wildcard - when the cost map carries its own azure_ai/ entry; the azure/ entry from the OpenAI-on-Azure - price sheet never reaches the Foundry provider list (LIT-7081).""" import litellm from litellm.proxy.auth.model_checks import get_known_models_from_wildcard diff --git a/tests/test_litellm/router_utils/test_reasoning_effort_capability.py b/tests/test_litellm/router_utils/test_reasoning_effort_capability.py index b7499d1c975..3e1f26b6e1c 100644 --- a/tests/test_litellm/router_utils/test_reasoning_effort_capability.py +++ b/tests/test_litellm/router_utils/test_reasoning_effort_capability.py @@ -400,8 +400,7 @@ class TestGpt6AstraAdvertisesItsDocumentedLevels: def test_a_foundry_deployment_also_advertises_none(self, local_model_cost_map, model, custom_llm_provider): """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, whether - it is reached through the azure route or the azure_ai (Foundry) route.""" + OpenAI's rejects, so an Azure deployment offers none on top of low through max.""" from litellm.utils import _get_model_info_helper model_info = dict(_get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider)) From e8f311429ea9afa09195bed21f078bbd50dd791e Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 5 Sep 2026 19:42:15 -0700 Subject: [PATCH 06/15] fix(cost-map): stop advertising reasoning_effort max on azure_ai/gpt-6-astra Foundry rejects reasoning_effort max on the gpt-6-astra deployment with a 400 that names none, low, medium, high, and xhigh as the supported values, so the card no longer lists max. The request path never gated max (only xhigh is opt-in), so this only changes /model_group/info and router capability gating. The azure/ twin stays as is because it was not verified on an Azure OpenAI host --- .../model_prices_and_context_window_backup.json | 2 +- model_prices_and_context_window.json | 2 +- .../test_reasoning_effort_capability.py | 14 +++++++++++++- 3 files changed, 15 insertions(+), 3 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 5b4652d38c7..b659c3b65e5 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -3518,7 +3518,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, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 5b4652d38c7..b659c3b65e5 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -3518,7 +3518,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, diff --git a/tests/test_litellm/router_utils/test_reasoning_effort_capability.py b/tests/test_litellm/router_utils/test_reasoning_effort_capability.py index 3e1f26b6e1c..fa3a6dcd95a 100644 --- a/tests/test_litellm/router_utils/test_reasoning_effort_capability.py +++ b/tests/test_litellm/router_utils/test_reasoning_effort_capability.py @@ -394,7 +394,6 @@ class TestGpt6AstraAdvertisesItsDocumentedLevels: [ ("azure/gpt-6-astra", "azure"), ("azure/us/gpt-6-astra", "azure"), - ("azure_ai/gpt-6-astra", "azure_ai"), ], ) def test_a_foundry_deployment_also_advertises_none(self, local_model_cost_map, model, custom_llm_provider): @@ -413,3 +412,16 @@ class TestGpt6AstraAdvertisesItsDocumentedLevels: "xhigh", "max", ) + + def test_a_foundry_azure_ai_deployment_advertises_none_but_not_max(self, local_model_cost_map): + from litellm.utils import _get_model_info_helper + + model_info = dict(_get_model_info_helper(model="azure_ai/gpt-6-astra", custom_llm_provider="azure_ai")) + + assert resolve_supported_reasoning_efforts(model_info, deployment_is_mapped=True) == ( + "none", + "low", + "medium", + "high", + "xhigh", + ) From 061c25b5cac15212ded745c51e3299aa4d43f056 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 5 Sep 2026 21:25:03 -0700 Subject: [PATCH 07/15] fix(spend): let a batch's charge survive an older proxy's $0 poll row A proxy running the old code wrote _batch_cost at $0 every time it polled a batch that was still running, so after an upgrade the claim found that row and read it as proof the batch had already been charged. Only a row that recorded a charge counts now, which leaves those $0 rows, and any row a client planted under the batch id, to be charged over disable_spend_logs skipped the claim entirely, so under that setting every retrieve of a finished batch charged again. The claim now runs either way and writes the one row per batch that makes the charge exactly once, while the per-request logs stay off --- litellm/proxy/db/db_spend_update_writer.py | 42 +++++++------ .../proxy/db/test_db_spend_update_writer.py | 60 ++++++++++++++++--- 2 files changed, 78 insertions(+), 24 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 3fad351224b..48312c025dd 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -271,9 +271,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: - if not await self._record_spend_log(payload=payload, prisma_client=prisma_client): - return False await self._enqueue_tool_usage_transaction( payload=payload, completion_response=completion_response, @@ -328,19 +331,23 @@ class DBSpendUpdateWriter: ) return True - async def _record_spend_log(self, payload: SpendLogsPayload, prisma_client: "PrismaClient | None") -> bool: - if prisma_client is None or not _is_batch_cost_row(payload): + 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) + if disable_spend_logs is False: await self._insert_spend_log_to_db(payload=payload, prisma_client=prisma_client) - return True - return await self._claim_batch_cost_spend_log(payload=payload, prisma_client=prisma_client) + return True async def _claim_batch_cost_spend_log(self, payload: SpendLogsPayload, prisma_client: "PrismaClient") -> 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 a successful retrieve wrote counts: a failed retrieve, or any request whose - client picked the batch id as its call id, cannot take the charge away. + 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 @@ -362,17 +369,18 @@ class DBSpendUpdateWriter: ) await self._insert_spend_log_to_db(payload=payload, prisma_client=prisma_client) return True - if ( - existing is not None - and existing.call_type == CallTypes.aretrieve_batch.value - and existing.status == "success" - ): + 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 - 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), + verbose_proxy_logger.debug( + "Spend row %s charged nothing for this batch, so this retrieve charges it", request_id ) return True diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index 500a0e7bb06..41f2be08545 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -2977,10 +2977,12 @@ def _spend_logs_prisma(inserted: int, existing: object) -> MagicMock: return prisma -async def _update_database_with(db_writer: DBSpendUpdateWriter, prisma: MagicMock, payload: dict) -> bool: +async def _update_database_with( + db_writer: DBSpendUpdateWriter, prisma: MagicMock, payload: dict, disable_spend_logs: bool = False +) -> 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", False + "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 @@ -3014,14 +3016,16 @@ async def _update_database_with(db_writer: DBSpendUpdateWriter, prisma: MagicMoc ("inserted", "existing", "charged"), [ (1, None, True), - (0, SimpleNamespace(call_type="aretrieve_batch", status="success"), False), - (0, SimpleNamespace(call_type="aretrieve_batch", status="failure"), True), - (0, SimpleNamespace(call_type="aembedding", status="success"), 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", @@ -3033,8 +3037,9 @@ async def test_update_database_charges_a_batch_only_from_the_retrieve_that_wrote """ 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 - written by anything but a successful retrieve, say a request whose client picked the - batch id as its call id, must not be able to take the charge away. + 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() @@ -3049,6 +3054,47 @@ async def test_update_database_charges_a_batch_only_from_the_retrieve_that_wrote assert db_writer._batch_database_updates.await_count == (1 if charged else 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 + + @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.""" From e79f3ec5205d01323093534a6577905a6dfdb7ac Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 5 Sep 2026 22:31:32 -0700 Subject: [PATCH 08/15] fix(cost-map): stop advertising reasoning_effort max on the azure gpt-6-astra rows Both Azure routes refuse it. A live call to the same deployment through openai/deployments/gpt-6-astra/chat/completions on api-version 2025-04-01-preview answers reasoning_effort max with a 400 unsupported_value naming none, low, medium, high and xhigh as the values it takes, and xhigh returns 200, so azure/gpt-6-astra and azure/us/gpt-6-astra now match the azure_ai row. --- ...odel_prices_and_context_window_backup.json | 4 +-- .../reasoning_effort_capability.py | 4 +-- model_prices_and_context_window.json | 4 +-- .../test_reasoning_effort_capability.py | 26 ++++++------------- 4 files changed, 14 insertions(+), 24 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index b659c3b65e5..48c51f8bbf5 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -7237,7 +7237,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, @@ -7503,7 +7503,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, diff --git a/litellm/router_utils/reasoning_effort_capability.py b/litellm/router_utils/reasoning_effort_capability.py index 9185d901a28..7b145c15a07 100644 --- a/litellm/router_utils/reasoning_effort_capability.py +++ b/litellm/router_utils/reasoning_effort_capability.py @@ -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. diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index b659c3b65e5..48c51f8bbf5 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -7237,7 +7237,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, @@ -7503,7 +7503,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, diff --git a/tests/test_litellm/router_utils/test_reasoning_effort_capability.py b/tests/test_litellm/router_utils/test_reasoning_effort_capability.py index fa3a6dcd95a..ccd6766b13a 100644 --- a/tests/test_litellm/router_utils/test_reasoning_effort_capability.py +++ b/tests/test_litellm/router_utils/test_reasoning_effort_capability.py @@ -394,30 +394,20 @@ class TestGpt6AstraAdvertisesItsDocumentedLevels: [ ("azure/gpt-6-astra", "azure"), ("azure/us/gpt-6-astra", "azure"), + ("azure_ai/gpt-6-astra", "azure_ai"), ], ) - def test_a_foundry_deployment_also_advertises_none(self, local_model_cost_map, model, custom_llm_provider): - """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.""" + 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=custom_llm_provider)) - assert resolve_supported_reasoning_efforts(model_info, deployment_is_mapped=True) == ( - "none", - "low", - "medium", - "high", - "xhigh", - "max", - ) - - def test_a_foundry_azure_ai_deployment_advertises_none_but_not_max(self, local_model_cost_map): - from litellm.utils import _get_model_info_helper - - model_info = dict(_get_model_info_helper(model="azure_ai/gpt-6-astra", custom_llm_provider="azure_ai")) - assert resolve_supported_reasoning_efforts(model_info, deployment_is_mapped=True) == ( "none", "low", From fa2b64878b6f7be8fed5139ef95961fe26241f99 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 5 Sep 2026 22:31:33 -0700 Subject: [PATCH 09/15] fix(azure_ai): redirect a gpt-5 capability lookup only when the map has a foundry row gpt-6-astra is the only gpt-5-family name with an azure_ai row. Prefixing the rest cost them every effort flag, since get_llm_provider sends an azure_ai name down the azure provider when a global AZURE_AI_API_BASE points at an openai.azure.com host and azure/ is not a key either, which turned temperature, top_p and logprobs on azure_ai/gpt-5.1-chat-latest from accepted into an UnsupportedParamsError. --- litellm/llms/azure_ai/chat/transformation.py | 14 ++++++++++- .../chat/test_azure_ai_transformation.py | 23 +++++++++++++++++++ 2 files changed, 36 insertions(+), 1 deletion(-) diff --git a/litellm/llms/azure_ai/chat/transformation.py b/litellm/llms/azure_ai/chat/transformation.py index 7c9a26c3f07..039c462b38a 100644 --- a/litellm/llms/azure_ai/chat/transformation.py +++ b/litellm/llms/azure_ai/chat/transformation.py @@ -46,7 +46,19 @@ NON_OPENAI_SPEC_MESSAGE_FIELDS: Final = ( class AzureAIGPT5Config(OpenAIGPT5Config): @classmethod def _model_map_lookup_name(cls, model: str) -> str: - return model if model.startswith("azure_ai/") else f"azure_ai/{model}" + """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/`` 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/`` 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() diff --git a/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py b/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py index 5ff0b729449..9924d77eb39 100644 --- a/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py +++ b/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py @@ -157,6 +157,29 @@ def test_foundry_gpt_6_astra_keeps_sampling_params_when_reasoning_effort_is_none 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 +): + """gpt-6-astra is the only gpt-5-family name with an azure_ai/ row. Reading an azure_ai/ key for + the rest finds nothing, and an openai.azure.com base sends that 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. From 0fb3951b2cadca5d9091b7586d0a5a15e4600b42 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 5 Sep 2026 22:32:06 -0700 Subject: [PATCH 10/15] fix(spend): charge a batch once when an older proxy left its cost row at $0 A proxy without this fix wrote the batch's cost row on every poll while the batch was still running, so that row reads $0 and the insert that claims the charge has nowhere to land. The retrieve that charges the batch now writes its own payload over that row under a where clause that still names spend 0.0, so exactly one retrieve takes it over and every later one reads the charge and charges nothing --- litellm/proxy/db/db_spend_update_writer.py | 42 +++++++++- .../proxy/db/test_db_spend_update_writer.py | 76 ++++++++++++++++++- 2 files changed, 112 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 48312c025dd..ae3bc5663eb 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -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 @@ -379,9 +380,44 @@ class DBSpendUpdateWriter: if existing.spend > 0: verbose_proxy_logger.debug("Cost tracking skipped: spend row %s already charged this batch", request_id) return False - verbose_proxy_logger.debug( - "Spend row %s charged nothing for this batch, so this retrieve charges it", request_id - ) + return await self._take_over_uncharged_batch_cost_row(payload=payload, prisma_client=prisma_client) + + async def _take_over_uncharged_batch_cost_row( + self, payload: SpendLogsPayload, prisma_client: "PrismaClient" + ) -> 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 payload.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; a row it cannot take over charges the batch + verbose_proxy_logger.warning("Could not take over spend row %s for a batch's cost: %s", request_id, e) + return True + 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( diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index 41f2be08545..33b7af06e1a 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -2969,16 +2969,21 @@ def _batch_cost_payload() -> dict: } -def _spend_logs_prisma(inserted: int, existing: object) -> MagicMock: +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 + 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 @@ -3005,7 +3010,7 @@ async def _update_database_with( completion_response=None, start_time=datetime.now(timezone.utc), end_time=datetime.now(timezone.utc), - response_cost=0.25, + response_cost=response_cost, ) await asyncio.sleep(0) return charged @@ -3054,6 +3059,71 @@ async def test_update_database_charges_a_batch_only_from_the_retrieve_that_wrote 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_charges_a_batch_whose_zero_row_it_could_not_take_over(): + """A DB that refuses the takeover must not swallow the batch's cost.""" + 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 True + + assert db_writer._batch_database_updates.await_count == 1 + + +@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"), From 24f0be80219cd3403ec28cddbed9be946fad1cb0 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 5 Sep 2026 22:47:41 -0700 Subject: [PATCH 11/15] fix(spend): leave a batch uncharged when the database refuses the takeover The takeover of a $0 row an older proxy left behind used to charge the batch when the update could not reach the database. That leaves the row still reading $0, so every later retrieve finds the same row and charges the batch again, which is the repeat charging this PR exists to stop. The retrieve that does take the row over is the one that charges, and a batch nobody retrieves again after that failure is never charged, the same as one whose proxy died inside the write window. --- litellm/proxy/db/db_spend_update_writer.py | 8 +++++--- .../proxy/db/test_db_spend_update_writer.py | 12 ++++++++---- 2 files changed, 13 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index ae3bc5663eb..ee7802a45d3 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -412,9 +412,11 @@ class DBSpendUpdateWriter: "spend": 0.0, }, ) - except Exception as e: # noqa: BLE001 # prisma raises its own hierarchy; a row it cannot take over charges the batch - verbose_proxy_logger.warning("Could not take over spend row %s for a batch's cost: %s", request_id, e) - return True + 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 diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index 33b7af06e1a..4efb94b60aa 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -3093,17 +3093,21 @@ async def test_update_database_charges_a_batch_whose_row_a_pre_upgrade_poll_left @pytest.mark.asyncio -async def test_update_database_charges_a_batch_whose_zero_row_it_could_not_take_over(): - """A DB that refuses the takeover must not swallow the batch's cost.""" +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 True + assert await _update_database_with(db_writer, prisma, _batch_cost_payload()) is False - assert db_writer._batch_database_updates.await_count == 1 + assert db_writer._batch_database_updates.await_count == 0 @pytest.mark.asyncio From 3dea1ebb32c96560f9f10171d0597ca30d4bea40 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 5 Sep 2026 23:18:25 -0700 Subject: [PATCH 12/15] fix(cost-map): keep the prompt cache breakpoint flag on the foundry gpt-6-astra row The openai gpt-6-astra card carries supports_prompt_cache_breakpoint, so a Foundry deployment reported it as true until the azure_ai row took over the lookup. The cache control hook still honours breakpoints for that deployment through the bare name, so /model/info was the only thing that changed, and it now agrees with the hook again. --- litellm/model_prices_and_context_window_backup.json | 1 + model_prices_and_context_window.json | 1 + 2 files changed, 2 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 48c51f8bbf5..2f86deffe54 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -3524,6 +3524,7 @@ "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, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 48c51f8bbf5..2f86deffe54 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -3524,6 +3524,7 @@ "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, From fffe0bb0dc94f8b071be85050a0bdaa101c0e2ba Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 5 Sep 2026 23:18:25 -0700 Subject: [PATCH 13/15] test(azure_ai): pin the tier the messages bridge sends when astra refuses max The /v1/messages adapter lowers a tier the entry does not accept, so dropping max from the astra rows moves that path from Foundry's 400 to a request at xhigh. Nothing pinned that, and the guard test's docstring named gpt-6-astra as the only gpt-5 name with an azure_ai row, which 11 rows contradict. --- ..._handler_reasoning_effort_normalization.py | 19 +++++++++++++++++++ .../chat/test_azure_ai_transformation.py | 8 ++++---- 2 files changed, 23 insertions(+), 4 deletions(-) diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_handler_reasoning_effort_normalization.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_handler_reasoning_effort_normalization.py index 56b754c3476..af7befecc33 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_handler_reasoning_effort_normalization.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_handler_reasoning_effort_normalization.py @@ -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" diff --git a/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py b/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py index 9924d77eb39..f8cc0b5071e 100644 --- a/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py +++ b/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py @@ -160,10 +160,10 @@ def test_foundry_gpt_6_astra_keeps_sampling_params_when_reasoning_effort_is_none def test_a_gpt_5_name_without_a_foundry_row_keeps_reading_its_own_entry( monkeypatch: pytest.MonkeyPatch, _local_model_cost_map ): - """gpt-6-astra is the only gpt-5-family name with an azure_ai/ row. Reading an azure_ai/ key for - the rest finds nothing, and an openai.azure.com base sends that 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.""" + """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") From fcb6d2267c09389ae1aa9e80e04a35caa5b8b470 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 5 Sep 2026 23:51:44 -0700 Subject: [PATCH 14/15] fix(spend): keep a batch's claim row out of the logs a proxy was told not to write disable_spend_logs has to keep meaning that no request gets logged, and the row that makes a batch chargeable exactly once is the one row it cannot drop, so with logging off that row now carries only what tells the retrieves apart. SPEND_LOGS_URL deployments get their copy back too: the claim writes straight to this table, so the row is queued as well when an external writer is the one that takes the spend logs. --- litellm/proxy/db/db_spend_update_writer.py | 49 +++++++-- .../proxy/db/test_db_spend_update_writer.py | 99 +++++++++++++++++++ 2 files changed, 141 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index ee7802a45d3..bdc014d7f13 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -89,6 +89,21 @@ 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 @@ -336,12 +351,16 @@ class DBSpendUpdateWriter: 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) + 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") -> bool: + 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 @@ -353,13 +372,17 @@ class DBSpendUpdateWriter: 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(payload)], # mutable-ok: prisma create_many takes a list + data=[prisma_client.jsonify_object(row)], # mutable-ok: prisma create_many takes a list skip_duplicates=True, ) if claimed == 1: + await self._forward_batch_cost_row( + row=row, prisma_client=prisma_client, disable_spend_logs=disable_spend_logs + ) return True existing: Final = await spend_logs.find_unique( where={"request_id": request_id} # mutable-ok: prisma where clause @@ -368,7 +391,7 @@ class DBSpendUpdateWriter: 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=payload, prisma_client=prisma_client) + 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( @@ -380,10 +403,22 @@ class DBSpendUpdateWriter: 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) + return await self._take_over_uncharged_batch_cost_row(payload=payload, prisma_client=prisma_client, row=row) + + async def _forward_batch_cost_row( + self, row: Mapping[str, object], prisma_client: "PrismaClient", disable_spend_logs: bool + ) -> None: + """Queue the claimed row for an external spend log writer, which the claim went around. + + With ``SPEND_LOGS_URL`` set the queue posts every spend log to that writer instead of + inserting it, so a batch's cost row reaches it only by being queued here as well. + """ + if disable_spend_logs is True or os.getenv("SPEND_LOGS_URL") is None: + return + await self._insert_spend_log_to_db(payload=prisma_client.jsonify_object(row), prisma_client=prisma_client) async def _take_over_uncharged_batch_cost_row( - self, payload: SpendLogsPayload, prisma_client: "PrismaClient" + 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. @@ -403,7 +438,7 @@ class DBSpendUpdateWriter: try: taken_over: Final = await SpendLogsRepository(prisma_client).table.update_many( data=prisma_client.jsonify_object( - MappingProxyType({field: value for field, value in payload.items() if field != "request_id"}) + MappingProxyType({field: value for field, value in row.items() if field != "request_id"}) ), where={ # mutable-ok: prisma where clause "request_id": request_id, diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index 4efb94b60aa..a8aeaced55f 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -1,6 +1,7 @@ import asyncio import copy import json +import os import re @@ -3169,6 +3170,104 @@ async def test_update_database_writes_no_ordinary_spend_row_with_spend_logs_disa 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 +@pytest.mark.parametrize( + ("spend_logs_url", "forwarded"), + [("http://spend-logs.internal", True), (None, False)], + ids=["an_external_writer_takes_the_rows", "rows_are_written_to_this_db"], +) +async def test_update_database_sends_a_claimed_batch_cost_row_on_to_an_external_spend_log_writer( + monkeypatch, spend_logs_url: str | None, forwarded: bool +): + """ + SPEND_LOGS_URL makes the flush post spend logs to that writer instead of inserting them, + and the claim writes straight to this table, so the batch's row reaches the writer only + by being queued as well. Queueing it with no writer configured would insert it twice. + """ + db_writer = DBSpendUpdateWriter() + db_writer._batch_database_updates = AsyncMock() + prisma = _spend_logs_prisma(1, None) + if spend_logs_url is None: + monkeypatch.delenv("SPEND_LOGS_URL", raising=False) + else: + monkeypatch.setenv("SPEND_LOGS_URL", spend_logs_url) + + assert await _update_database_with(db_writer, prisma, _batch_cost_payload()) is True + + queued = [row["request_id"] for row in prisma.spend_log_transactions] + assert queued == (["batch_abc_batch_cost"] if forwarded else []) + + +@pytest.mark.asyncio +async def test_update_database_forwards_no_batch_cost_row_a_later_retrieve_had_already_claimed(monkeypatch): + """The retrieve that lost the claim charges nothing, so it must not post a row either.""" + db_writer = DBSpendUpdateWriter() + db_writer._batch_database_updates = AsyncMock() + existing = SimpleNamespace(call_type="aretrieve_batch", status="success", spend=0.25) + prisma = _spend_logs_prisma(0, existing) + monkeypatch.setenv("SPEND_LOGS_URL", "http://spend-logs.internal") + + assert await _update_database_with(db_writer, prisma, _batch_cost_payload()) is False + + assert prisma.spend_log_transactions == [] + + +@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.""" From defd8661f4e359994bf57a7f9e51ed8c479f17ff Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sun, 6 Sep 2026 00:04:24 -0700 Subject: [PATCH 15/15] refactor(spend): stop queueing a batch's claim row for a writer the proxy never builds SPEND_LOGS_URL only diverts spend logs when db_writer_client is set, and nothing in the proxy ever assigns that global, so the queued copy was only ever skipped as a duplicate by the local insert. --- litellm/proxy/db/db_spend_update_writer.py | 15 ------- .../proxy/db/test_db_spend_update_writer.py | 43 ------------------- 2 files changed, 58 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index bdc014d7f13..9230be8055e 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -380,9 +380,6 @@ class DBSpendUpdateWriter: skip_duplicates=True, ) if claimed == 1: - await self._forward_batch_cost_row( - row=row, prisma_client=prisma_client, disable_spend_logs=disable_spend_logs - ) return True existing: Final = await spend_logs.find_unique( where={"request_id": request_id} # mutable-ok: prisma where clause @@ -405,18 +402,6 @@ class DBSpendUpdateWriter: return False return await self._take_over_uncharged_batch_cost_row(payload=payload, prisma_client=prisma_client, row=row) - async def _forward_batch_cost_row( - self, row: Mapping[str, object], prisma_client: "PrismaClient", disable_spend_logs: bool - ) -> None: - """Queue the claimed row for an external spend log writer, which the claim went around. - - With ``SPEND_LOGS_URL`` set the queue posts every spend log to that writer instead of - inserting it, so a batch's cost row reaches it only by being queued here as well. - """ - if disable_spend_logs is True or os.getenv("SPEND_LOGS_URL") is None: - return - await self._insert_spend_log_to_db(payload=prisma_client.jsonify_object(row), prisma_client=prisma_client) - async def _take_over_uncharged_batch_cost_row( self, payload: SpendLogsPayload, prisma_client: "PrismaClient", row: Mapping[str, object] ) -> bool: diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index a8aeaced55f..0bca7c9492c 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -1,7 +1,6 @@ import asyncio import copy import json -import os import re @@ -3212,48 +3211,6 @@ async def test_update_database_claims_a_batch_without_logging_the_request_that_p assert db_writer._batch_database_updates.await_count == 1 -@pytest.mark.asyncio -@pytest.mark.parametrize( - ("spend_logs_url", "forwarded"), - [("http://spend-logs.internal", True), (None, False)], - ids=["an_external_writer_takes_the_rows", "rows_are_written_to_this_db"], -) -async def test_update_database_sends_a_claimed_batch_cost_row_on_to_an_external_spend_log_writer( - monkeypatch, spend_logs_url: str | None, forwarded: bool -): - """ - SPEND_LOGS_URL makes the flush post spend logs to that writer instead of inserting them, - and the claim writes straight to this table, so the batch's row reaches the writer only - by being queued as well. Queueing it with no writer configured would insert it twice. - """ - db_writer = DBSpendUpdateWriter() - db_writer._batch_database_updates = AsyncMock() - prisma = _spend_logs_prisma(1, None) - if spend_logs_url is None: - monkeypatch.delenv("SPEND_LOGS_URL", raising=False) - else: - monkeypatch.setenv("SPEND_LOGS_URL", spend_logs_url) - - assert await _update_database_with(db_writer, prisma, _batch_cost_payload()) is True - - queued = [row["request_id"] for row in prisma.spend_log_transactions] - assert queued == (["batch_abc_batch_cost"] if forwarded else []) - - -@pytest.mark.asyncio -async def test_update_database_forwards_no_batch_cost_row_a_later_retrieve_had_already_claimed(monkeypatch): - """The retrieve that lost the claim charges nothing, so it must not post a row either.""" - db_writer = DBSpendUpdateWriter() - db_writer._batch_database_updates = AsyncMock() - existing = SimpleNamespace(call_type="aretrieve_batch", status="success", spend=0.25) - prisma = _spend_logs_prisma(0, existing) - monkeypatch.setenv("SPEND_LOGS_URL", "http://spend-logs.internal") - - assert await _update_database_with(db_writer, prisma, _batch_cost_payload()) is False - - assert prisma.spend_log_transactions == [] - - @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."""