diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py index ec00e5cd03a..cd52fa6e9ec 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py @@ -1,8 +1,7 @@ """ Polls LiteLLM_ManagedObjectTable to check if the response is complete. -Cost tracking is handled by the get-responses call, which prices normally only because the -poll stamps itself with BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN; user-facing reads of the -same route are non-inference and free. +The status CAS is the durable billed marker and precedes the billing read, so a stolen or +duplicate claim can never bill the same row twice. """ from datetime import datetime, timedelta, timezone @@ -110,10 +109,10 @@ class CheckResponsesCost: async def _claim_job_for_costing(self, job: "LiteLLM_ManagedObjectTable") -> bool: """Atomically flip batch_processed false to true, returning whether this pod won the row. - Every pod polls the same table and the read is what prices the job, so the claim has to be - taken before it. The ``updated_at`` arm takes a claim back from a pod that died holding it; - ``updated_at`` is ``@updatedAt``, so a live claim is never stolen. A schema without the - column cannot claim, so it keeps the pre-existing behavior instead of billing nothing. + The claim precedes the probe and billing reads. The ``updated_at`` arm takes a claim back + from a pod that died holding it; ``updated_at`` is ``@updatedAt``, so a live claim is never + stolen. A schema without the column cannot claim, so it keeps the pre-existing behavior + instead of billing nothing. """ abandoned_before: Final = datetime.now(timezone.utc) - timedelta( seconds=CLAIM_ABANDONED_AFTER_POLL_CYCLES * PROXY_BATCH_POLLING_INTERVAL @@ -155,26 +154,31 @@ class CheckResponsesCost: f"so its cost will not be retried: {db_err}" ) - async def _mark_job_completed(self, job: "LiteLLM_ManagedObjectTable") -> None: - """Retire a billed row from polling, per job so one failure can't strand the rest. + async def _mark_job_completed(self, job: "LiteLLM_ManagedObjectTable") -> bool: + """CAS the durable billed marker before the billing read, returning whether this pod won. - Only ``status`` is written, and always the literal "completed", because the usage already - landed in ``LiteLLM_SpendLogs`` and stale-row expiry keys off that exact value. + Only queued or in-progress rows may transition to the literal "completed" value. """ try: - await self.prisma_client.db.litellm_managedobjecttable.update_many( - where={"id": job.id}, + updated: Final = await self.prisma_client.db.litellm_managedobjecttable.update_many( + where={ + "id": job.id, + "status": {"in": ["queued", "in_progress"]}, + }, data={"status": "completed"}, ) except Exception as db_err: verbose_proxy_logger.error( f"CheckResponsesCost: failed to mark job {job.id} completed: {db_err}" ) + return False + return updated > 0 async def check_responses_cost(self): """Read every queued background response and retire the ones the provider has finished. - The read itself is what bills, because it is stamped with the poll's call origin. + The probe is free. The status CAS marks the row billed before the follow-up read records + its spend. """ try: await self._cleanup_stale_managed_objects() @@ -193,7 +197,7 @@ class CheckResponsesCost: ) verbose_proxy_logger.debug(f"Found {len(jobs)} response jobs to check") - completed_jobs = [] + completed_count: int = 0 for job in jobs: unified_object_id = job.unified_object_id @@ -210,19 +214,13 @@ class CheckResponsesCost: # Decrypts rows written before model_object_id held the provider's own id. responses_id_security = ResponsesIDSecurity().provider_response_id(job.model_object_id) - # Prepare metadata with model information for cost tracking - litellm_metadata = { + probe_metadata: Final = { "user_api_key_user_id": job.created_by or "default-user-id", - INTERNAL_CALL_ORIGIN_METADATA_KEY: BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN, **({"user_api_key_team_id": job.team_id} if job.team_id else {}), **({"user_api_key": job.api_key, "user_api_key_hash": job.api_key} if job.api_key else {}), + **({"model": model_name, "model_group": model_name} if model_name else {}), } - # Add model information if available - if model_name: - litellm_metadata["model"] = model_name - litellm_metadata["model_group"] = model_name # Use same value for model_group - except Exception as e: verbose_proxy_logger.warning( f"Skipping job {unified_object_id} due to error: {e}" @@ -238,7 +236,7 @@ class CheckResponsesCost: try: response = await self._get_response( response_id=responses_id_security, - litellm_metadata=litellm_metadata, + litellm_metadata=probe_metadata, ) except Exception as e: await self._release_job_claim(job) @@ -255,15 +253,29 @@ class CheckResponsesCost: await self._release_job_claim(job) continue - verbose_proxy_logger.info( - f"Response {unified_object_id} has terminal status {response.status}, marking as complete" - ) - completed_jobs.append(job) + if not await self._mark_job_completed(job): + verbose_proxy_logger.debug( + f"Response {unified_object_id} was already finalized by another poller" + ) + continue - for job in completed_jobs: - await self._mark_job_completed(job) - - if len(completed_jobs) > 0: + completed_count += 1 verbose_proxy_logger.info( - f"Marked {len(completed_jobs)} response jobs as completed" + f"Response {unified_object_id} has terminal status {response.status}, marked as complete" ) + billing_metadata: Final = { + **probe_metadata, + INTERNAL_CALL_ORIGIN_METADATA_KEY: BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN, + } + try: + await self._get_response( + response_id=responses_id_security, + litellm_metadata=billing_metadata, + ) + except Exception as e: + verbose_proxy_logger.error( + f"Response {unified_object_id} is already finalized, so its spend will not be retried: {e}" + ) + + if completed_count > 0: + verbose_proxy_logger.info(f"Marked {completed_count} response jobs as completed") diff --git a/tests/proxy_unit_tests/test_check_responses_cost.py b/tests/proxy_unit_tests/test_check_responses_cost.py index 4851669ff2a..e6475c8b452 100644 --- a/tests/proxy_unit_tests/test_check_responses_cost.py +++ b/tests/proxy_unit_tests/test_check_responses_cost.py @@ -2,9 +2,8 @@ Unit tests for CheckResponsesCost class """ -import asyncio from datetime import datetime, timedelta, timezone -from unittest.mock import AsyncMock, MagicMock, Mock, patch +from unittest.mock import AsyncMock, MagicMock, call, patch import pytest @@ -461,7 +460,13 @@ class TestCheckResponsesCost: # Run the check with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: - mock_aget.side_effect = [mock_response1, mock_response2, mock_response3] + mock_aget.side_effect = [ + mock_response1, + mock_response1, + mock_response2, + mock_response3, + mock_response3, + ] await check_responses_cost_instance.check_responses_cost() @@ -636,7 +641,7 @@ class TestCheckResponsesCost: mock_sdk_aget.return_value = mock_response await check_responses_cost_instance.check_responses_cost() - mock_sdk_aget.assert_called_once() + assert mock_sdk_aget.await_count == 2 mock_llm_router.aget_responses.assert_not_called() @pytest.mark.asyncio @@ -687,9 +692,13 @@ class TestCheckResponsesCost: mock_sdk_aget.return_value = mock_response await check_responses_cost_instance.check_responses_cost() - mock_llm_router.get_deployment.assert_called_once_with(model_id="deployment-deleted") + assert mock_llm_router.get_deployment.call_count == 2 + assert mock_llm_router.get_deployment.call_args_list == [ + call(model_id="deployment-deleted"), + call(model_id="deployment-deleted"), + ] mock_llm_router.aget_responses.assert_not_called() - mock_sdk_aget.assert_called_once() + assert mock_sdk_aget.await_count == 2 assert mock_sdk_aget.call_args[1]["response_id"] == encoded_response_id assert _completed_job_ids(mock_prisma_client) == ["job-missing-deployment"] @@ -799,12 +808,15 @@ class TestCheckResponsesCost: mock_aget.return_value = mock_response await check_responses_cost_instance.check_responses_cost() - metadata = mock_aget.call_args[1]["litellm_metadata"] - assert metadata[INTERNAL_CALL_ORIGIN_METADATA_KEY] == "background_response_cost_poll" - assert metadata["user_api_key_team_id"] == "team-billed" - assert metadata["user_api_key"] == "sk-billed" - assert metadata["user_api_key_hash"] == "sk-billed" - assert is_unbilled_non_inference_call("aget_responses", metadata) is False + probe_metadata = mock_aget.call_args_list[0][1]["litellm_metadata"] + billing_metadata = mock_aget.call_args_list[1][1]["litellm_metadata"] + assert INTERNAL_CALL_ORIGIN_METADATA_KEY not in probe_metadata + assert billing_metadata[INTERNAL_CALL_ORIGIN_METADATA_KEY] == "background_response_cost_poll" + assert billing_metadata["user_api_key_team_id"] == "team-billed" + assert billing_metadata["user_api_key"] == "sk-billed" + assert billing_metadata["user_api_key_hash"] == "sk-billed" + assert is_unbilled_non_inference_call("aget_responses", probe_metadata) is True + assert is_unbilled_non_inference_call("aget_responses", billing_metadata) is False assert is_unbilled_non_inference_call("aget_responses", None) is True @pytest.mark.asyncio @@ -845,11 +857,12 @@ class TestCheckResponsesCost: ): """A pod that dies holding a claim strands the row forever, so the lease has to outlast a live cycle and still fire well before stale expiry gives up on the row unbilled.""" - from litellm.constants import PROXY_BATCH_POLLING_INTERVAL from litellm_enterprise.proxy.common_utils.check_responses_cost import ( CLAIM_ABANDONED_AFTER_POLL_CYCLES, ) + from litellm.constants import PROXY_BATCH_POLLING_INTERVAL + mock_job = MagicMock() mock_job.unified_object_id = "resp_test_abandoned" mock_job.model_object_id = _routed_response_id("resp_test_abandoned") @@ -906,7 +919,7 @@ class TestCheckResponsesCost: writes_and_reads = [] async def record_update_many(**kwargs): - writes_and_reads.append(kwargs["data"]) + writes_and_reads.append(kwargs) return 1 async def record_read(**kwargs): @@ -930,10 +943,184 @@ class TestCheckResponsesCost: await check_responses_cost_instance.check_responses_cost() - assert len(writes_and_reads) == 3 - assert writes_and_reads[0] == {"batch_processed": True} + assert len(writes_and_reads) == 4 + assert writes_and_reads[0]["data"] == {"batch_processed": True} assert writes_and_reads[1] == "provider_read" - assert writes_and_reads[2]["status"] == "completed" + assert writes_and_reads[2]["data"] == {"status": "completed"} + assert writes_and_reads[2]["where"] == { + "id": "job-ordering", + "status": {"in": ["queued", "in_progress"]}, + } + assert writes_and_reads[3] == "provider_read" + + @pytest.mark.asyncio + async def test_probe_read_is_free_and_only_the_post_cas_read_is_billed( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY + from litellm.types.utils import BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN + + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_probe_billing" + mock_job.model_object_id = _routed_response_id("resp_test_probe_billing") + mock_job.created_by = "test-user" + mock_job.id = "job-probe-billing" + mock_job.file_object = {"model": "gpt-5", "id": "resp_test_probe_billing"} + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + + response = ResponsesAPIResponse( + id="resp_probe_billing", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage(input_tokens=10, output_tokens=5, total_tokens=15), + ) + call_order = [] + + async def record_update_many(**kwargs): + call_order.append(("update", kwargs)) + return 1 + + async def record_read(**kwargs): + call_order.append(("read", kwargs)) + return response + + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + side_effect=record_update_many + ) + mock_llm_router.aget_responses = AsyncMock(side_effect=record_read) + + await check_responses_cost_instance.check_responses_cost() + + assert mock_llm_router.aget_responses.await_count == 2 + assert [kind for kind, _ in call_order] == ["update", "read", "update", "read"] + probe_metadata = call_order[1][1]["litellm_metadata"] + billing_metadata = call_order[3][1]["litellm_metadata"] + assert INTERNAL_CALL_ORIGIN_METADATA_KEY not in probe_metadata + assert billing_metadata[INTERNAL_CALL_ORIGIN_METADATA_KEY] == ( + BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN + ) + assert call_order[2][1]["where"] == { + "id": "job-probe-billing", + "status": {"in": ["queued", "in_progress"]}, + } + + @pytest.mark.asyncio + async def test_row_already_finalized_by_another_poller_is_not_billed( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY + + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_finalized" + mock_job.model_object_id = _routed_response_id("resp_test_finalized") + mock_job.created_by = "test-user" + mock_job.id = "job-finalized" + mock_job.file_object = {"model": "gpt-5", "id": "resp_test_finalized"} + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + side_effect=[1, 0] + ) + mock_llm_router.aget_responses = AsyncMock( + return_value=ResponsesAPIResponse( + id="resp_finalized", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + ) + + await check_responses_cost_instance.check_responses_cost() + + assert mock_llm_router.aget_responses.await_count == 1 + metadata = mock_llm_router.aget_responses.call_args.kwargs["litellm_metadata"] + assert INTERNAL_CALL_ORIGIN_METADATA_KEY not in metadata + completion_calls = _completion_calls(mock_prisma_client) + assert len(completion_calls) == 1 + assert completion_calls[0].kwargs["where"]["status"] == { + "in": ["queued", "in_progress"] + } + assert _release_calls(mock_prisma_client) == [] + + @pytest.mark.asyncio + async def test_non_terminal_probe_does_not_finalize_or_bill( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_non_terminal" + mock_job.model_object_id = _routed_response_id("resp_test_non_terminal") + mock_job.created_by = "test-user" + mock_job.id = "job-non-terminal" + mock_job.file_object = {"model": "gpt-5", "id": "resp_test_non_terminal"} + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + return_value=1 + ) + mock_llm_router.aget_responses = AsyncMock( + return_value=ResponsesAPIResponse( + id="resp_non_terminal", + object="response", + status="queued", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + ) + + await check_responses_cost_instance.check_responses_cost() + + assert mock_llm_router.aget_responses.await_count == 1 + assert _completion_calls(mock_prisma_client) == [] + assert len(_release_calls(mock_prisma_client)) == 1 + + @pytest.mark.asyncio + async def test_billing_read_failure_after_cas_does_not_reopen_the_row( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_billing_failure" + mock_job.model_object_id = _routed_response_id("resp_test_billing_failure") + mock_job.created_by = "test-user" + mock_job.id = "job-billing-failure" + mock_job.file_object = {"model": "gpt-5", "id": "resp_test_billing_failure"} + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + return_value=1 + ) + mock_llm_router.aget_responses = AsyncMock( + side_effect=[ + ResponsesAPIResponse( + id="resp_billing_failure", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ), + Exception("boom"), + ] + ) + + await check_responses_cost_instance.check_responses_cost() + + assert mock_llm_router.aget_responses.await_count == 2 + assert len(_completion_calls(mock_prisma_client)) == 1 + assert _release_calls(mock_prisma_client) == [] + assert all( + call.kwargs["data"] not in ({"batch_processed": False}, {"status": "queued"}, {"status": "in_progress"}) + for call in mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list + ) @pytest.mark.asyncio @pytest.mark.parametrize("provider_status", ["queued", "in_progress"]) @@ -1049,7 +1236,7 @@ class TestCheckResponsesCost: await check_responses_cost_instance.check_responses_cost() - mock_llm_router.aget_responses.assert_awaited_once() + assert mock_llm_router.aget_responses.await_count == 2 assert mock_llm_router.aget_responses.await_args.kwargs["response_id"] == _routed_response_id( "resp_test_second" ) @@ -1140,15 +1327,14 @@ class TestCheckResponsesCost: await check_responses_cost_instance.check_responses_cost() - mock_llm_router.aget_responses.assert_awaited_once() + assert mock_llm_router.aget_responses.await_count == 2 assert _completed_job_ids(mock_prisma_client) == ["job-old-schema"] @pytest.mark.asyncio async def test_a_failed_persist_does_not_abort_the_rest_of_the_poll_cycle( self, check_responses_cost_instance, mock_prisma_client, mock_llm_router ): - """One row's write failing must not take the whole cycle down with it: the jobs behind it - are already read and billed, so losing their write loses their usage for good.""" + """One row's completion write failing must not take the whole cycle down with it.""" mock_job1 = MagicMock() mock_job1.unified_object_id = "resp_test_persist_fails" mock_job1.model_object_id = _routed_response_id("resp_test_persist_fails") @@ -1166,8 +1352,16 @@ class TestCheckResponsesCost: mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( return_value=[mock_job1, mock_job2] ) + async def fail_first_completion_write(**kwargs): + if ( + kwargs["data"] == {"status": "completed"} + and kwargs["where"]["id"] == "job-persist-fails" + ): + raise Exception("deadlock detected") + return 1 + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - side_effect=[1, 1, Exception("deadlock detected"), 1] + side_effect=fail_first_completion_write ) mock_response = ResponsesAPIResponse( @@ -1183,7 +1377,7 @@ class TestCheckResponsesCost: await check_responses_cost_instance.check_responses_cost() - assert mock_llm_router.aget_responses.await_count == 2 + assert mock_llm_router.aget_responses.await_count == 3 assert _completed_job_ids(mock_prisma_client) == [ "job-persist-fails", "job-persist-works",