From bb81a9f9f173747907a9a354af058e410a7679b0 Mon Sep 17 00:00:00 2001 From: jesus Date: Tue, 8 Sep 2026 15:06:37 +0000 Subject: [PATCH 01/13] fix(responses): bill a background response once via the cost poller, not on every read Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../common_utils/check_responses_cost.py | 3 +- .../proxy/hooks/managed_files.py | 10 +++---- litellm/integrations/opentelemetry.py | 10 ++----- litellm/integrations/otel/plumbing/metrics.py | 2 +- .../internal_call_metadata.py | 28 ++++--------------- litellm/litellm_core_utils/litellm_logging.py | 4 +-- litellm/proxy/common_request_processing.py | 2 +- .../proxy/response_api_endpoints/endpoints.py | 1 + .../spend_tracking/spend_tracking_utils.py | 2 +- .../test_check_responses_cost.py | 15 ++++++++-- .../integrations/otel/test_otel_v2_metrics.py | 12 +++----- .../integrations/test_opentelemetry.py | 12 ++++---- .../test_responses_background_cost.py | 2 ++ .../test_internal_call_metadata.py | 14 +++++++++- .../test_litellm_logging.py | 11 ++++---- .../test_spend_tracking_utils.py | 8 ++---- .../proxy/test_common_request_processing.py | 4 +-- 17 files changed, 68 insertions(+), 72 deletions(-) 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 06cf5fcf82f..7e569664c75 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py @@ -159,6 +159,8 @@ class CheckResponsesCost: litellm_metadata = { "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 {}), } # Add model information if available @@ -196,4 +198,3 @@ class CheckResponsesCost: verbose_proxy_logger.info( f"Marked {len(completed_jobs)} response jobs as completed" ) - diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index 4899b87da7a..39e1cb3a9a3 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -317,11 +317,11 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): ) -> None: """Persist a managed object row, caching it and upserting it in the DB. - persist_attribution is set only by the batch create, which is the one caller - that can speak for the creator; it gates the api_key and request_tags columns - that CheckBatchCost bills against, so a later poll or retrieve of the same - batch cannot record itself as the paying key. Like created_by and team_id, - both are written only in the upsert create branch, never on update. + persist_attribution is set by creates that can speak for the creator; it gates + the api_key and request_tags columns that cost pollers bill against, so a later + poll or retrieve of the same object cannot record itself as the paying key. + Like created_by and team_id, both are written only in the upsert create branch, + never on update. create_if_missing is cleared by callers that observe a batch they did not create, such as a poll. They still refresh status and file_object, but a diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index d4e7fcb577e..bd20b1dcc7d 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -1647,7 +1647,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): if ( self._token_usage_histogram and response_obj - and not is_unbilled_non_inference_call_from_params(kwargs.get("call_type"), params, response_obj) + and not is_unbilled_non_inference_call_from_params(kwargs.get("call_type"), params) and (usage := response_obj.get("usage")) ): in_attrs: Final = {**common_attrs, TOKEN_TYPE_ATTRIBUTE: "input"} @@ -1725,9 +1725,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): if not self._time_per_output_token_histogram: return - if is_unbilled_non_inference_call_from_params( - kwargs.get("call_type"), kwargs.get("litellm_params"), response_obj - ): + if is_unbilled_non_inference_call_from_params(kwargs.get("call_type"), kwargs.get("litellm_params")): return # Get completion tokens from response_obj @@ -2502,9 +2500,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): usage: Final = ( response_obj.get("usage") if response_obj - and not is_unbilled_non_inference_call_from_params( - kwargs.get("call_type"), litellm_params, response_obj - ) + and not is_unbilled_non_inference_call_from_params(kwargs.get("call_type"), litellm_params) else None ) if usage: diff --git a/litellm/integrations/otel/plumbing/metrics.py b/litellm/integrations/otel/plumbing/metrics.py index e1623f4697f..7bc080c36ee 100644 --- a/litellm/integrations/otel/plumbing/metrics.py +++ b/litellm/integrations/otel/plumbing/metrics.py @@ -224,7 +224,7 @@ class GenAIMetricRecorder: common_attrs: Final = self._filter_attributes(self._bounded_attributes(kwargs)) duration_s: Final = (end_time - start_time).total_seconds() usage_is_replayed: Final = is_unbilled_non_inference_call_from_params( - kwargs.get("call_type"), kwargs.get("litellm_params"), response_obj + kwargs.get("call_type"), kwargs.get("litellm_params") ) self._metrics.operation_duration.record(duration_s, attributes=common_attrs) diff --git a/litellm/litellm_core_utils/internal_call_metadata.py b/litellm/litellm_core_utils/internal_call_metadata.py index 87f007ca1d5..c4ab791a0d3 100644 --- a/litellm/litellm_core_utils/internal_call_metadata.py +++ b/litellm/litellm_core_utils/internal_call_metadata.py @@ -54,35 +54,18 @@ budget-checked like the request that spawned it. Everything else on the parent's be a lie on a sub-call that runs after it returned.""" -def is_background_response(response: object) -> bool: - """Whether a retrieved object is a response created with ``background=true``. - - Such a create returns ``status="queued"`` and no usage at all, so nothing has billed the - job by the time anyone reads it back. Accepts the response as a mapping or a model, - because the callers hold it in both shapes. - """ - if isinstance(response, Mapping): - return response.get("background") is True - return getattr(response, "background", None) is True - - def is_unbilled_non_inference_call( call_type: str | None, metadata: Mapping[str, object] | None, - response: object, ) -> bool: - """A read/management route priced at zero, because the usage it reports belongs to the - call that created the object it just read. + """Reads of stored objects are priced at zero because their usage belongs to the create. - Retrieving a background response is the exception, and the enterprise cost poller's read - is the same exception seen from the other side: that job's create billed nothing, so its - retrieval is the only place the spend is ever visible. Pricing those at zero would lose - the spend rather than deduplicate it. + A background create bills nothing, so the enterprise cost poller's read stamped with + ``BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN`` prices normally and marks the managed-object + row completed so it prices once. """ if call_type not in NON_INFERENCE_CALL_TYPES: return False - if is_background_response(response): - return False if metadata is None: return True return metadata.get(INTERNAL_CALL_ORIGIN_METADATA_KEY) != BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN @@ -91,7 +74,6 @@ def is_unbilled_non_inference_call( def is_unbilled_non_inference_call_from_params( call_type: str | None, litellm_params: Mapping[str, object] | None, - response: object, ) -> bool: """:func:`is_unbilled_non_inference_call` for callers holding raw ``litellm_params``. @@ -105,7 +87,7 @@ def is_unbilled_non_inference_call_from_params( metadata: Final = ( StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params) if litellm_params is not None else None ) - return is_unbilled_non_inference_call(call_type, metadata, response) + return is_unbilled_non_inference_call(call_type, metadata) def sanitize_user_api_key_auth(auth: object) -> object: diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 9ba9fd082f3..1bf2ad6c4ab 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -1720,7 +1720,7 @@ class Logging(LiteLLMLoggingBaseClass): return 0.0 if is_unbilled_non_inference_call( - self.call_type, StandardLoggingPayloadSetup.merge_litellm_metadata(self.litellm_params), result + self.call_type, StandardLoggingPayloadSetup.merge_litellm_metadata(self.litellm_params) ): return 0.0 @@ -6153,7 +6153,7 @@ def get_standard_logging_object_payload( cache_hit: Final = kwargs.get("cache_hit", False) # Extract usage as a plain dict, avoiding Pydantic round-trip raw_usage_dict: Final = StandardLoggingPayloadSetup.get_usage_as_dict( - response_obj=None if is_unbilled_non_inference_call(call_type, metadata, response_obj) else response_obj, + response_obj=None if is_unbilled_non_inference_call(call_type, metadata) else response_obj, combined_usage_object=cast(Usage | None, kwargs.get("combined_usage_object")), ) usage_dict: Final = ( diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 0ad86479aac..bde6c6bccb6 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -2694,7 +2694,7 @@ class ProxyBaseLLMRequestProcessing: ) llm_cost_for_headers: Final = ( 0.0 - if is_unbilled_non_inference_call_from_params(logging_obj.call_type, logging_obj.litellm_params, response) + if is_unbilled_non_inference_call_from_params(logging_obj.call_type, logging_obj.litellm_params) else computed_cost_for_headers ) _, request_metadata_bucket = get_or_create_metadata_bucket(self.data) diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index 5907ffc64eb..f320e2c65fb 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -397,6 +397,7 @@ async def responses_api( model_object_id=response.id, file_purpose="response", user_api_key_dict=user_api_key_dict, + persist_attribution=True, ) verbose_proxy_logger.info( diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index 4bcdf6aad22..49cbbc0184b 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -371,7 +371,7 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs usage: dict = {} if call_type in ["ocr", "aocr"]: usage = _extract_usage_for_ocr_call(response_obj, response_obj_dict) - elif not is_unbilled_non_inference_call(call_type, metadata, response_obj_dict): + elif not is_unbilled_non_inference_call(call_type, metadata): # Use response_obj_dict instead of response_obj to avoid calling .get() on Pydantic models _usage: Final = response_obj_dict.get("usage", None) or {} if isinstance(_usage, litellm.Usage): diff --git a/tests/proxy_unit_tests/test_check_responses_cost.py b/tests/proxy_unit_tests/test_check_responses_cost.py index e806e9a3394..b206891e5e9 100644 --- a/tests/proxy_unit_tests/test_check_responses_cost.py +++ b/tests/proxy_unit_tests/test_check_responses_cost.py @@ -732,6 +732,8 @@ class TestCheckResponsesCost: mock_job = MagicMock() mock_job.unified_object_id = "resp_test_no_model" mock_job.created_by = "test-user" + mock_job.team_id = None + mock_job.api_key = None mock_job.id = "job-no-model" mock_job.file_object = {} # no "model" key → model_name=None branch @@ -753,6 +755,9 @@ class TestCheckResponsesCost: call_kwargs = mock_aget.call_args[1] assert "model" not in call_kwargs.get("litellm_metadata", {}) assert "model_group" not in call_kwargs.get("litellm_metadata", {}) + assert "user_api_key_team_id" not in call_kwargs["litellm_metadata"] + assert "user_api_key" not in call_kwargs["litellm_metadata"] + assert "user_api_key_hash" not in call_kwargs["litellm_metadata"] @pytest.mark.asyncio async def test_poll_stamps_internal_call_origin_so_the_read_is_billed( @@ -769,6 +774,8 @@ class TestCheckResponsesCost: mock_job = MagicMock() mock_job.unified_object_id = "resp_test_billed" mock_job.created_by = "test-user" + mock_job.team_id = "team-billed" + mock_job.api_key = "sk-billed" mock_job.id = "job-billed" mock_job.file_object = {"model": "gpt-5", "id": "resp_test_billed"} @@ -787,7 +794,9 @@ class TestCheckResponsesCost: await check_responses_cost_instance.check_responses_cost() metadata = mock_aget.call_args[1]["litellm_metadata"] - foreground_read = {"background": False} assert metadata[INTERNAL_CALL_ORIGIN_METADATA_KEY] == "background_response_cost_poll" - assert is_unbilled_non_inference_call("aget_responses", metadata, foreground_read) is False - assert is_unbilled_non_inference_call("aget_responses", None, foreground_read) is True + 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 + assert is_unbilled_non_inference_call("aget_responses", None) is True diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_metrics.py b/tests/test_litellm/integrations/otel/test_otel_v2_metrics.py index 016dbcd824b..32f1bd8ceda 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_metrics.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_metrics.py @@ -214,10 +214,8 @@ def test_response_read_does_not_replay_the_generation_usage(): assert RESPONSE_DURATION in metrics -def test_background_response_read_still_records_usage(): - """A background=true create returns no usage, so its completed read is the only - place the generation's tokens are ever seen. Skipping it would lose them - entirely rather than deduplicate them.""" +def test_background_response_read_does_not_record_usage(): + """The enterprise cost poller owns usage metrics for completed background responses.""" reader = InMemoryMetricReader() logger = _logger(reader, enable_metrics=True) kwargs, response_obj, start, end = _build_call(call_type="aget_responses") @@ -225,10 +223,8 @@ def test_background_response_read_still_records_usage(): asyncio.run(logger.async_log_success_event(kwargs, response_obj, start, end)) metrics = _metrics_by_name(reader) - by_type = {dp.attributes[TOKEN_TYPE]: dp for dp in metrics[TOKEN_USAGE]} - assert by_type["input"].sum == PROMPT_TOKENS - assert by_type["output"].sum == COMPLETION_TOKENS - assert TIME_PER_OUTPUT_TOKEN in metrics + assert TOKEN_USAGE not in metrics + assert TIME_PER_OUTPUT_TOKEN not in metrics def test_metrics_disabled_records_nothing(): diff --git a/tests/test_litellm/integrations/test_opentelemetry.py b/tests/test_litellm/integrations/test_opentelemetry.py index 9ec8489f784..96388c6fa12 100644 --- a/tests/test_litellm/integrations/test_opentelemetry.py +++ b/tests/test_litellm/integrations/test_opentelemetry.py @@ -6418,14 +6418,14 @@ class TestOpenTelemetryNonInferenceUsage(unittest.TestCase): def test_background_cost_poll_read_still_records_the_token_usage_histogram(self): self.assertEqual(self._token_histogram_calls("aget_responses", self.BACKGROUND_POLL), 2) - def test_background_response_read_still_reports_its_tokens_on_the_span(self): + def test_background_response_read_does_not_report_its_tokens_on_the_span(self): self.assertEqual( self._token_attributes_on_span("aget_responses", response_obj=self.BACKGROUND_RESPONSE_OBJ), - set(self.TOKEN_KEYS), + set(), ) - def test_background_response_read_still_records_the_token_usage_histogram(self): - self.assertEqual(self._token_histogram_calls("aget_responses", response_obj=self.BACKGROUND_RESPONSE_OBJ), 2) + def test_background_response_read_does_not_record_the_token_usage_histogram(self): + self.assertEqual(self._token_histogram_calls("aget_responses", response_obj=self.BACKGROUND_RESPONSE_OBJ), 0) def test_inference_call_still_records_time_per_output_token(self): self.assertEqual(self._time_per_output_token_calls("acompletion"), 1) @@ -6433,7 +6433,7 @@ class TestOpenTelemetryNonInferenceUsage(unittest.TestCase): def test_response_read_does_not_divide_its_latency_by_the_retrieved_token_count(self): self.assertEqual(self._time_per_output_token_calls("aget_responses"), 0) - def test_background_response_read_still_records_time_per_output_token(self): + def test_background_response_read_does_not_record_time_per_output_token(self): self.assertEqual( - self._time_per_output_token_calls("aget_responses", response_obj=self.BACKGROUND_RESPONSE_OBJ), 1 + self._time_per_output_token_calls("aget_responses", response_obj=self.BACKGROUND_RESPONSE_OBJ), 0 ) diff --git a/tests/test_litellm/integrations/test_responses_background_cost.py b/tests/test_litellm/integrations/test_responses_background_cost.py index 0d4218f2137..7ad39194222 100644 --- a/tests/test_litellm/integrations/test_responses_background_cost.py +++ b/tests/test_litellm/integrations/test_responses_background_cost.py @@ -81,6 +81,7 @@ class TestResponsesBackgroundCostTracking: model_object_id=response.id, file_purpose="response", user_api_key_dict=user_api_key_dict, + persist_attribution=True, ) # Verify store_unified_object_id was called @@ -92,6 +93,7 @@ class TestResponsesBackgroundCostTracking: assert call_args[1]["model_object_id"] == response.id assert call_args[1]["file_purpose"] == "response" assert call_args[1]["user_api_key_dict"] == user_api_key_dict + assert call_args[1]["persist_attribution"] is True @pytest.mark.asyncio async def test_no_storage_for_non_background_requests( diff --git a/tests/test_litellm/litellm_core_utils/test_internal_call_metadata.py b/tests/test_litellm/litellm_core_utils/test_internal_call_metadata.py index 73923dc75a5..e69e5313aae 100644 --- a/tests/test_litellm/litellm_core_utils/test_internal_call_metadata.py +++ b/tests/test_litellm/litellm_core_utils/test_internal_call_metadata.py @@ -3,9 +3,10 @@ from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY from litellm.litellm_core_utils.internal_call_metadata import ( forwarded_internal_call_metadata, + is_unbilled_non_inference_call, sanitized_forwardable_call_metadata, ) -from litellm.types.utils import SHADOW_EVAL_ROUTER_CALL_ORIGIN +from litellm.types.utils import BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN, SHADOW_EVAL_ROUTER_CALL_ORIGIN PARENT = { "user_api_key": "sk-hash", @@ -49,6 +50,17 @@ def test_sanitized_forwardable_metadata_keeps_only_identity_and_always_stamps(): } +def test_background_response_reads_are_free_but_cost_poller_reads_are_billed(): + assert is_unbilled_non_inference_call("aget_responses", {}) is True + assert ( + is_unbilled_non_inference_call( + "aget_responses", + {INTERNAL_CALL_ORIGIN_METADATA_KEY: BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN}, + ) + is False + ) + + class TestSubCallMetadataSanitization: """The proxy cost callback must not be able to recover the parent budget reservation from sub-call metadata, in either of the shapes it knows how to read.""" 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 70f9bae283b..c4faf9efc6f 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -5622,15 +5622,14 @@ class TestNonInferenceCallTypesAreNotBilled: assert payload is not None assert payload["total_tokens"] == 6000 - def test_reading_a_background_response_is_still_priced(self): - """A background create answers queued with no usage at all, so whoever reads the finished - job is the first and only caller to see its tokens. Zeroing that read bills the job nothing.""" + def test_reading_a_background_response_is_free(self): + """A completed background response read is free because the cost poller owns its billing.""" cost = self._logging_obj("aget_responses")._response_cost_calculator( result=self._retrieved_response(background=True) ) - assert cost is not None and cost > 0 + assert cost == 0.0 - def test_reading_a_background_response_reports_usage_in_standard_logging_payload(self): + def test_reading_a_background_response_does_not_report_usage_in_standard_logging_payload(self): from datetime import datetime from litellm.litellm_core_utils.litellm_logging import ( @@ -5653,7 +5652,7 @@ class TestNonInferenceCallTypesAreNotBilled: ) assert payload is not None - assert payload["total_tokens"] == 6000 + assert payload["total_tokens"] == 0 def test_reading_a_foreground_response_is_still_free(self): """Guards the test above against a blanket exemption: an explicit background=false read was diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py index a72b4e28143..9d1018023fa 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py @@ -3612,13 +3612,11 @@ def test_spend_log_for_background_response_cost_poll_counts_tokens(): assert payload["total_tokens"] == 6000 -def test_spend_log_for_background_response_retrieval_counts_tokens(): - """A background create answers queued carrying no usage, so its retrieval is the first and only - place the job's tokens are ever visible. Zeroing that read bills the whole job nothing on any - proxy that is not running the enterprise cost poller.""" +def test_spend_log_for_background_response_retrieval_does_not_count_tokens(): + """The enterprise cost poller owns billing for background response usage.""" payload = _spend_log_for_call_type("aget_responses", background=True) - assert payload["total_tokens"] == 6000 + assert payload["total_tokens"] == 0 def test_spend_log_for_foreground_response_retrieval_still_counts_nothing(): diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 812fd8ed47d..27c2f6bd10f 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -5848,7 +5848,7 @@ class TestCostHeadersForCallsPricedAtZero: assert fastapi_response.headers[f"x-litellm-response-cost-{component}"] == "0.0" @pytest.mark.asyncio - async def test_reading_a_background_response_keeps_its_real_cost(self, monkeypatch): + async def test_reading_a_background_response_keeps_its_zero_cost(self, monkeypatch): fastapi_response = await self._drive( monkeypatch=monkeypatch, response=self._responses_read(background=True), @@ -5856,7 +5856,7 @@ class TestCostHeadersForCallsPricedAtZero: route_type="aget_responses", ) - assert float(fastapi_response.headers["x-litellm-response-cost"]) == pytest.approx(0.00042) + assert float(fastapi_response.headers["x-litellm-response-cost"]) == 0.0 @pytest.mark.asyncio async def test_an_inference_call_without_a_recorded_cost_still_omits_the_header(self, monkeypatch): From 0235dbd7f2932815b01e680f8d5bc9657c7079d4 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Wed, 9 Sep 2026 15:11:45 -0700 Subject: [PATCH 02/13] fix(responses): claim a background response before the read that bills it Every pod and uvicorn worker schedules its own CheckResponsesCost against the shared LiteLLM_ManagedObjectTable. The poller selected eligible rows, performed the billed retrieval, and only then marked them completed in one bulk write, so two pollers could select the same terminal response and both record a charge before either completion update landed. Each row is now claimed with a compare-and-swap on batch_processed before the read, because the read is what prices the job: aget_responses stamped with the poll origin writes the spend log itself, so there is no later point at which to serialize. A row whose read raised, or whose provider status is still non-terminal, releases its claim so a later cycle retries it rather than retiring it unbilled. That is the failure #37050 fixed on the batch side. A pod that dies between winning the claim and billing would otherwise strand the row: it holds a claim nobody will release and its status never reaches terminal, so every later cycle re-selects it and loses. The updated_at arm of the claim takes such a row back after three poll cycles, and since updated_at is @updatedAt a healthy in-flight claim written moments ago is never stolen. The poller now also persists the finished response onto its managed row instead of writing status alone, so the row carries the generation's usage rather than the stale queued copy stored at create time. Reuses the existing batch_processed column, so no migration. It already sits on the shared table defaulted to false and was unused by response rows. Claude-Session: https://claude.ai/code/session_01Hn5E8Jz1LjGLFyiYxBRcBW --- .../common_utils/check_responses_cost.py | 142 ++++- .../test_check_responses_cost.py | 571 +++++++++++++++--- .../test_responses_background_cost.py | 34 +- 3 files changed, 635 insertions(+), 112 deletions(-) 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 7e569664c75..62e742be680 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py @@ -6,7 +6,7 @@ same route are non-inference and free. """ from datetime import datetime, timedelta, timezone -from typing import TYPE_CHECKING, Dict, Optional, cast +from typing import TYPE_CHECKING, Dict, Final, Optional, cast import litellm from litellm._logging import verbose_proxy_logger @@ -14,6 +14,7 @@ from litellm.constants import ( INTERNAL_CALL_ORIGIN_METADATA_KEY, MANAGED_OBJECT_STALENESS_CUTOFF_DAYS, MAX_OBJECTS_PER_POLL_CYCLE, + PROXY_BATCH_POLLING_INTERVAL, STALE_OBJECT_CLEANUP_BATCH_SIZE, ) from litellm.responses.utils import ResponsesAPIRequestUtils @@ -21,11 +22,14 @@ from litellm.types.llms.openai import ResponsesAPIResponse from litellm.types.utils import BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN if TYPE_CHECKING: + from litellm.proxy._types import LiteLLM_ManagedObjectTable from litellm.proxy.utils import PrismaClient, ProxyLogging from litellm.router import Router TERMINAL_RESPONSE_STATUSES = frozenset({"completed", "failed", "cancelled", "incomplete"}) +CLAIM_ABANDONED_AFTER_POLL_CYCLES: Final = 3 + class CheckResponsesCost: def __init__( @@ -112,6 +116,92 @@ class CheckResponsesCost: f"(older than {MANAGED_OBJECT_STALENESS_CUTOFF_DAYS} days) as stale_expired" ) + @staticmethod + def _is_missing_batch_processed_column_error(err: Exception) -> bool: + message: Final = str(err).lower() + return "batch_processed" in message or "unknown column" in message or "does not exist" in message + + async def _claim_job_for_costing(self, job: "LiteLLM_ManagedObjectTable") -> bool: + """Atomically flip batch_processed from false to true, returning whether this pod won the row. + + Every pod and uvicorn worker schedules its own CheckResponsesCost against the shared table, + so without this compare-and-swap two of them select the same queued response in one window + and both bill it. The claim is taken before the read because the read is what prices the + job: ``aget_responses`` stamped with the poll origin writes the spend log itself, so there + is no later point at which to serialize. Schemas without the column can't be claimed, so + they keep the pre-existing behavior rather than silently billing nothing. + + A pod that dies between winning the claim and billing would otherwise strand the row: + it holds a claim nobody will release, and its status never reaches terminal, so every + later cycle re-selects it and loses. The ``updated_at`` arm takes such a claim back once + it has gone unbilled for longer than any live cycle could hold it. ``updated_at`` is + ``@updatedAt``, so a healthy in-flight claim refreshed moments ago is never stolen. + """ + abandoned_before: Final = datetime.now(timezone.utc) - timedelta( + seconds=CLAIM_ABANDONED_AFTER_POLL_CYCLES * PROXY_BATCH_POLLING_INTERVAL + ) + try: + claimed: Final = await self.prisma_client.db.litellm_managedobjecttable.update_many( + where={ + "id": job.id, + "OR": [ + {"batch_processed": False}, + {"updated_at": {"lt": abandoned_before}}, + ], + }, + data={"batch_processed": True}, + ) + except Exception as db_err: + if self._is_missing_batch_processed_column_error(db_err): + verbose_proxy_logger.warning( + "CheckResponsesCost: batch_processed column not found, billing without a claim" + ) + return True + verbose_proxy_logger.error(f"CheckResponsesCost: failed to claim job {job.id} for cost tracking: {db_err}") + return False + return claimed > 0 + + async def _release_job_claim(self, job: "LiteLLM_ManagedObjectTable") -> None: + """Give a claimed row back when the read did not bill it, so a later poll cycle retries it. + + A response still queued at the provider, or whose read raised, has no spend to record yet. + Holding the claim would retire it permanently, which is the failure #37050 hit on batches. + """ + try: + await self.prisma_client.db.litellm_managedobjecttable.update_many( + where={"id": job.id, "batch_processed": True}, + data={"batch_processed": False}, + ) + except Exception as db_err: + verbose_proxy_logger.error( + f"CheckResponsesCost: failed to release the claim on job {job.id}, " + f"so its cost will not be retried: {db_err}" + ) + + async def _persist_terminal_response( + self, job: "LiteLLM_ManagedObjectTable", response: ResponsesAPIResponse + ) -> None: + """Store the finished response on its managed row and retire the row from polling. + + The row is the only copy of a background generation's usage that outlives the poll, so + ``GET /v1/responses/{id}`` can serve a terminal job from here instead of re-reading it + from the provider. Every provider re-read replays the same usage and hands back a + freshly encoded id, which is what made this route bill per read and defeated id-based + dedup in the first place. + + ``status`` stays the literal "completed" for every terminal provider status, matching + what this poller has always written, so stale-row expiry keeps skipping these rows. + """ + try: + await self.prisma_client.db.litellm_managedobjecttable.update_many( + where={"id": job.id}, + data={"status": "completed", "file_object": response.model_dump_json()}, + ) + except Exception as db_err: + verbose_proxy_logger.error( + f"CheckResponsesCost: failed to persist terminal response for job {job.id}: {db_err}" + ) + async def check_responses_cost(self): """ Check if background responses are complete and track their cost. @@ -168,33 +258,47 @@ class CheckResponsesCost: litellm_metadata["model"] = model_name litellm_metadata["model_group"] = model_name # Use same value for model_group - response = await self._get_response( - response_id=responses_id_security, - litellm_metadata=litellm_metadata, - ) - - verbose_proxy_logger.debug( - f"Response {unified_object_id} status: {response.status}, model: {model_name}" - ) - except Exception as e: verbose_proxy_logger.warning( f"Skipping job {unified_object_id} due to error: {e}" ) continue - if response.status in TERMINAL_RESPONSE_STATUSES: - verbose_proxy_logger.info( - f"Response {unified_object_id} has terminal status {response.status}, marking as complete" + if not await self._claim_job_for_costing(job): + verbose_proxy_logger.debug( + f"Response {unified_object_id} is already claimed for costing, leaving it to the claim holder" ) - completed_jobs.append(job) + continue - # Mark completed jobs in the database - if len(completed_jobs) > 0: - await self.prisma_client.db.litellm_managedobjecttable.update_many( - where={"id": {"in": [job.id for job in completed_jobs]}}, - data={"status": "completed"}, + try: + response = await self._get_response( + response_id=responses_id_security, + litellm_metadata=litellm_metadata, + ) + except Exception as e: + await self._release_job_claim(job) + verbose_proxy_logger.warning( + f"Skipping job {unified_object_id} due to error: {e}" + ) + continue + + verbose_proxy_logger.debug( + f"Response {unified_object_id} status: {response.status}, model: {model_name}" ) + + if response.status not in TERMINAL_RESPONSE_STATUSES: + 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, response)) + + for job, response in completed_jobs: + await self._persist_terminal_response(job, response) + + if len(completed_jobs) > 0: verbose_proxy_logger.info( f"Marked {len(completed_jobs)} 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 b206891e5e9..2d832975a11 100644 --- a/tests/proxy_unit_tests/test_check_responses_cost.py +++ b/tests/proxy_unit_tests/test_check_responses_cost.py @@ -3,7 +3,8 @@ Unit tests for CheckResponsesCost class """ import asyncio -from datetime import datetime +import json +from datetime import datetime, timedelta, timezone from unittest.mock import AsyncMock, MagicMock, Mock, patch import pytest @@ -12,6 +13,40 @@ from litellm.constants import MAX_OBJECTS_PER_POLL_CYCLE from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse +def _update_many_calls_writing(mock_prisma_client, matches_data): + return [ + call + for call in mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list + if matches_data(call.kwargs["data"]) + ] + + +def _completion_calls(mock_prisma_client): + return _update_many_calls_writing( + mock_prisma_client, lambda data: data.get("status") == "completed" + ) + + +def _completed_job_ids(mock_prisma_client): + return [call.kwargs["where"]["id"] for call in _completion_calls(mock_prisma_client)] + + +def _persisted_response(completion_call): + return json.loads(completion_call.kwargs["data"]["file_object"]) + + +def _claim_calls(mock_prisma_client): + return _update_many_calls_writing( + mock_prisma_client, lambda data: data == {"batch_processed": True} + ) + + +def _release_calls(mock_prisma_client): + return _update_many_calls_writing( + mock_prisma_client, lambda data: data == {"batch_processed": False} + ) + + class TestCheckResponsesCost: """Test suite for CheckResponsesCost class""" @@ -135,7 +170,7 @@ class TestCheckResponsesCost: ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) # Run the check with mocked litellm.aget_responses @@ -144,14 +179,8 @@ class TestCheckResponsesCost: await check_responses_cost_instance.check_responses_cost() - # update_many should only contain the job completion call - calls = ( - mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list - ) - assert len(calls) == 1 - completion_call = calls[0] - assert completion_call[1]["data"]["status"] == "completed" - assert completion_call[1]["where"]["id"]["in"] == ["job-123"] + assert _completed_job_ids(mock_prisma_client) == ["job-123"] + assert _release_calls(mock_prisma_client) == [] @pytest.mark.asyncio async def test_check_responses_cost_with_failed_response( @@ -180,7 +209,7 @@ class TestCheckResponsesCost: ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) # Run the check @@ -189,12 +218,8 @@ class TestCheckResponsesCost: await check_responses_cost_instance.check_responses_cost() - # update_many should only contain the job completion call - calls = ( - mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list - ) - assert len(calls) == 1 - assert calls[0][1]["data"]["status"] == "completed" + assert _completed_job_ids(mock_prisma_client) == ["job-456"] + assert _release_calls(mock_prisma_client) == [] @pytest.mark.asyncio async def test_check_responses_cost_with_cancelled_response( @@ -223,7 +248,7 @@ class TestCheckResponsesCost: ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) # Run the check @@ -232,12 +257,8 @@ class TestCheckResponsesCost: await check_responses_cost_instance.check_responses_cost() - # update_many should only contain the job completion call - calls = ( - mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list - ) - assert len(calls) == 1 - assert calls[0][1]["data"]["status"] == "completed" + assert _completed_job_ids(mock_prisma_client) == ["job-789"] + assert _release_calls(mock_prisma_client) == [] @pytest.mark.asyncio async def test_check_responses_cost_with_in_progress_response( @@ -266,7 +287,7 @@ class TestCheckResponsesCost: ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) # Run the check @@ -276,10 +297,7 @@ class TestCheckResponsesCost: await check_responses_cost_instance.check_responses_cost() # No job completion update_many — response is still in progress - calls = ( - mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list - ) - assert len(calls) == 0 + assert _completion_calls(mock_prisma_client) == [] # Stale cleanup still ran via _expire_stale_rows check_responses_cost_instance._expire_stale_rows.assert_called_once() @@ -310,7 +328,7 @@ class TestCheckResponsesCost: ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) # Run the check @@ -320,10 +338,7 @@ class TestCheckResponsesCost: await check_responses_cost_instance.check_responses_cost() # No job completion update_many — response is still queued - calls = ( - mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list - ) - assert len(calls) == 0 + assert _completion_calls(mock_prisma_client) == [] # Stale cleanup still ran via _expire_stale_rows check_responses_cost_instance._expire_stale_rows.assert_called_once() @@ -344,7 +359,7 @@ class TestCheckResponsesCost: ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) # Run the check with mocked exception @@ -357,10 +372,7 @@ class TestCheckResponsesCost: await check_responses_cost_instance.check_responses_cost() # No job completion update_many — exception skipped the job - calls = ( - mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list - ) - assert len(calls) == 0 + assert _completion_calls(mock_prisma_client) == [] # Stale cleanup still ran via _expire_stale_rows check_responses_cost_instance._expire_stale_rows.assert_called_once() @@ -429,7 +441,7 @@ class TestCheckResponsesCost: ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) # Run the check @@ -438,16 +450,7 @@ class TestCheckResponsesCost: await check_responses_cost_instance.check_responses_cost() - # update_many should only contain the job completion call - calls = ( - mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list - ) - assert len(calls) == 1 - completion_call = calls[0] - assert len(completion_call[1]["where"]["id"]["in"]) == 2 - assert "job-1" in completion_call[1]["where"]["id"]["in"] - assert "job-3" in completion_call[1]["where"]["id"]["in"] - assert "job-2" not in completion_call[1]["where"]["id"]["in"] + assert _completed_job_ids(mock_prisma_client) == ["job-1", "job-3"] @pytest.mark.asyncio async def test_encoded_response_id_is_fetched_through_router( @@ -480,7 +483,7 @@ class TestCheckResponsesCost: return_value=[mock_job] ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) mock_llm_router.aget_responses = AsyncMock( @@ -511,12 +514,7 @@ class TestCheckResponsesCost: == encoded_response_id ) - calls = ( - mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list - ) - assert len(calls) == 1 - assert calls[0][1]["data"]["status"] == "completed" - assert calls[0][1]["where"]["id"]["in"] == ["job-router"] + assert _completed_job_ids(mock_prisma_client) == ["job-router"] @pytest.mark.asyncio async def test_encrypted_response_id_is_fetched_through_router( @@ -556,7 +554,7 @@ class TestCheckResponsesCost: return_value=[mock_job] ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) mock_llm_router.aget_responses = AsyncMock( @@ -584,11 +582,7 @@ class TestCheckResponsesCost: mock_llm_router.aget_responses.call_args[1]["response_id"] == encoded_response_id ) - calls = ( - mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list - ) - assert len(calls) == 1 - assert calls[0][1]["where"]["id"]["in"] == ["job-encrypted"] + assert _completed_job_ids(mock_prisma_client) == ["job-encrypted"] @pytest.mark.asyncio async def test_response_id_without_model_id_uses_sdk( @@ -605,7 +599,7 @@ class TestCheckResponsesCost: return_value=[mock_job] ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) mock_llm_router.aget_responses = AsyncMock( side_effect=AssertionError("router cannot route an id without a model_id") @@ -654,7 +648,7 @@ class TestCheckResponsesCost: return_value=[mock_job] ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) mock_llm_router.get_deployment = MagicMock(return_value=None) mock_llm_router.aget_responses = AsyncMock( @@ -679,12 +673,7 @@ class TestCheckResponsesCost: mock_sdk_aget.assert_called_once() assert mock_sdk_aget.call_args[1]["response_id"] == encoded_response_id - calls = ( - mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list - ) - assert len(calls) == 1 - assert calls[0][1]["data"]["status"] == "completed" - assert calls[0][1]["where"]["id"]["in"] == ["job-missing-deployment"] + assert _completed_job_ids(mock_prisma_client) == ["job-missing-deployment"] @pytest.mark.asyncio async def test_check_responses_cost_with_incomplete_response( @@ -701,7 +690,7 @@ class TestCheckResponsesCost: return_value=[mock_job] ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) mock_response = ResponsesAPIResponse( @@ -717,12 +706,7 @@ class TestCheckResponsesCost: mock_aget.return_value = mock_response await check_responses_cost_instance.check_responses_cost() - calls = ( - mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list - ) - assert len(calls) == 1 - assert calls[0][1]["data"]["status"] == "completed" - assert calls[0][1]["where"]["id"]["in"] == ["job-incomplete"] + assert _completed_job_ids(mock_prisma_client) == ["job-incomplete"] @pytest.mark.asyncio async def test_check_responses_cost_no_model_in_file_object( @@ -741,7 +725,7 @@ class TestCheckResponsesCost: return_value=[mock_job] ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) mock_response = MagicMock() @@ -783,7 +767,7 @@ class TestCheckResponsesCost: return_value=[mock_job] ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) mock_response = MagicMock() @@ -800,3 +784,430 @@ class TestCheckResponsesCost: assert metadata["user_api_key_hash"] == "sk-billed" assert is_unbilled_non_inference_call("aget_responses", metadata) is False assert is_unbilled_non_inference_call("aget_responses", None) is True + + @pytest.mark.asyncio + async def test_job_claimed_by_another_pod_is_never_read_or_completed( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + """Every pod and uvicorn worker polls the same table, and the read is what writes the + spend log, so losing the claim has to skip the read entirely or the job is billed twice.""" + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_claimed_elsewhere" + mock_job.created_by = "test-user" + mock_job.id = "job-claimed-elsewhere" + mock_job.file_object = {"model": "gpt-5", "id": "resp_test_claimed_elsewhere"} + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + return_value=0 + ) + + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_sdk_aget: + await check_responses_cost_instance.check_responses_cost() + + mock_sdk_aget.assert_not_awaited() + mock_llm_router.aget_responses.assert_not_awaited() + assert _completion_calls(mock_prisma_client) == [] + assert _release_calls(mock_prisma_client) == [] + + claim_calls = _claim_calls(mock_prisma_client) + assert len(claim_calls) == 1 + claim_where = claim_calls[0].kwargs["where"] + assert claim_where["id"] == "job-claimed-elsewhere" + assert {"batch_processed": False} in claim_where["OR"] + + @pytest.mark.asyncio + async def test_claim_is_taken_back_from_a_pod_that_died_holding_it( + self, check_responses_cost_instance, mock_prisma_client + ): + """A pod that dies between claiming and billing releases nothing, and the row's status + never reaches terminal, so without a lease every later cycle re-selects it and loses. + The window has to be longer than a live cycle can hold a claim and short enough that the + row is retried well before stale expiry gives up on it unbilled.""" + from litellm.constants import PROXY_BATCH_POLLING_INTERVAL + from litellm_enterprise.proxy.common_utils.check_responses_cost import ( + CLAIM_ABANDONED_AFTER_POLL_CYCLES, + ) + + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_abandoned" + mock_job.created_by = "test-user" + mock_job.id = "job-abandoned" + mock_job.file_object = {"model": "gpt-5", "id": "resp_test_abandoned"} + + 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_response = ResponsesAPIResponse( + id="resp_abandoned", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage(input_tokens=100, output_tokens=50, total_tokens=150), + ) + + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: + mock_aget.return_value = mock_response + await check_responses_cost_instance.check_responses_cost() + + claim_where = _claim_calls(mock_prisma_client)[0].kwargs["where"] + abandoned_arm = next(arm for arm in claim_where["OR"] if "updated_at" in arm) + lease = timedelta( + seconds=CLAIM_ABANDONED_AFTER_POLL_CYCLES * PROXY_BATCH_POLLING_INTERVAL + ) + untouched_for = datetime.now(timezone.utc) - abandoned_arm["updated_at"]["lt"] + assert lease <= untouched_for < lease + timedelta(seconds=30) + assert lease > timedelta(seconds=PROXY_BATCH_POLLING_INTERVAL) + + @pytest.mark.asyncio + async def test_claim_is_taken_before_the_billing_read_and_kept_on_a_terminal_status( + self, check_responses_cost_instance, mock_prisma_client + ): + """The read prices the job, so the claim has to be taken before it, and keeping the claim + afterwards is what stops a second pod reading and billing the same row again.""" + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_ordering" + mock_job.created_by = "test-user" + mock_job.id = "job-ordering" + mock_job.file_object = {"model": "gpt-5", "id": "resp_test_ordering"} + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + + writes_and_reads = [] + + async def record_update_many(**kwargs): + writes_and_reads.append(kwargs["data"]) + return 1 + + async def record_read(**kwargs): + writes_and_reads.append("provider_read") + return ResponsesAPIResponse( + id="resp_ordering", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage( + input_tokens=100, output_tokens=50, total_tokens=150 + ), + ) + + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + side_effect=record_update_many + ) + + with patch( + "litellm.aget_responses", new_callable=AsyncMock, side_effect=record_read + ): + await check_responses_cost_instance.check_responses_cost() + + assert len(writes_and_reads) == 3 + assert writes_and_reads[0] == {"batch_processed": True} + assert writes_and_reads[1] == "provider_read" + assert writes_and_reads[2]["status"] == "completed" + + @pytest.mark.asyncio + @pytest.mark.parametrize("provider_status", ["queued", "in_progress"]) + async def test_non_terminal_status_releases_the_claim( + self, check_responses_cost_instance, mock_prisma_client, provider_status + ): + """A response the provider has not finished yet has no spend to record, so its row must go + back to batch_processed=False; holding the claim retires it before it is ever billed.""" + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_still_running" + mock_job.created_by = "test-user" + mock_job.id = "job-still-running" + mock_job.file_object = {"model": "gpt-5", "id": "resp_test_still_running"} + + 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_response = ResponsesAPIResponse( + id="resp_still_running", + object="response", + status=provider_status, + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: + mock_aget.return_value = mock_response + await check_responses_cost_instance.check_responses_cost() + + assert _completion_calls(mock_prisma_client) == [] + release_calls = _release_calls(mock_prisma_client) + assert len(release_calls) == 1 + assert release_calls[0].kwargs["where"] == { + "id": "job-still-running", + "batch_processed": True, + } + + @pytest.mark.asyncio + async def test_failed_provider_read_releases_the_claim( + self, check_responses_cost_instance, mock_prisma_client + ): + """A read that raised billed nothing, so the claim has to be handed back or the row is + retired unbilled and no later poll cycle ever retries it.""" + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_read_error" + mock_job.created_by = "test-user" + mock_job.id = "job-read-error" + mock_job.file_object = {"model": "gpt-5", "id": "resp_test_read_error"} + + 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 + ) + + with patch( + "litellm.aget_responses", + new_callable=AsyncMock, + side_effect=Exception("Provider error"), + ): + await check_responses_cost_instance.check_responses_cost() + + assert _completion_calls(mock_prisma_client) == [] + release_calls = _release_calls(mock_prisma_client) + assert len(release_calls) == 1 + assert release_calls[0].kwargs["where"] == { + "id": "job-read-error", + "batch_processed": True, + } + + @pytest.mark.asyncio + async def test_a_job_claimed_elsewhere_does_not_block_the_next_job( + self, check_responses_cost_instance, mock_prisma_client + ): + """Losing one row to another pod must skip only that row: the rest of the poll page still + has to be read and billed in the same cycle.""" + mock_job1 = MagicMock() + mock_job1.unified_object_id = "resp_test_first" + mock_job1.created_by = "user1" + mock_job1.id = "job-first" + mock_job1.file_object = {"model": "gpt-5", "id": "resp_test_first"} + + mock_job2 = MagicMock() + mock_job2.unified_object_id = "resp_test_second" + mock_job2.created_by = "user2" + mock_job2.id = "job-second" + mock_job2.file_object = {"model": "gpt-5", "id": "resp_test_second"} + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job1, mock_job2] + ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + side_effect=[0, 1, 1] + ) + + mock_response = ResponsesAPIResponse( + id="resp_second", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage(input_tokens=100, output_tokens=50, total_tokens=150), + ) + + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: + mock_aget.return_value = mock_response + await check_responses_cost_instance.check_responses_cost() + + mock_aget.assert_awaited_once() + assert mock_aget.await_args.kwargs["response_id"] == "resp_test_second" + + assert _completed_job_ids(mock_prisma_client) == ["job-second"] + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "db_error_message", + [ + "column LiteLLM_ManagedObjectTable.batch_processed does not exist", + "Unknown column in where clause", + "The column P2022 does not exist in the current database", + ], + ) + async def test_claim_fails_open_on_a_schema_without_the_claim_column( + self, check_responses_cost_instance, mock_prisma_client, db_error_message + ): + """A deployment that never ran the batch_processed migration cannot claim anything, so it + keeps the pre-claim behavior of billing rather than silently billing nothing.""" + mock_job = MagicMock() + mock_job.id = "job-old-schema" + + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + side_effect=Exception(db_error_message) + ) + + assert ( + await check_responses_cost_instance._claim_job_for_costing(mock_job) is True + ) + + @pytest.mark.asyncio + async def test_claim_is_lost_when_the_database_fails_for_any_other_reason( + self, check_responses_cost_instance, mock_prisma_client + ): + """A dropped connection is no proof the row is free, so the read that would bill it is + not allowed to run.""" + mock_job = MagicMock() + mock_job.id = "job-db-down" + + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + side_effect=Exception("connection to server was lost") + ) + + assert ( + await check_responses_cost_instance._claim_job_for_costing(mock_job) is False + ) + + @pytest.mark.asyncio + async def test_old_schema_without_the_claim_column_still_bills_and_completes( + self, check_responses_cost_instance, mock_prisma_client + ): + """End to end on a pre-migration schema: the claim write fails, the response is still read + (which is what bills it) and the row is still marked completed.""" + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_old_schema" + mock_job.created_by = "test-user" + mock_job.id = "job-old-schema" + mock_job.file_object = {"model": "gpt-5", "id": "resp_test_old_schema"} + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + + async def reject_batch_processed_writes(**kwargs): + if "batch_processed" in kwargs["data"]: + raise Exception( + 'column "batch_processed" of relation ' + '"LiteLLM_ManagedObjectTable" does not exist' + ) + return 1 + + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + side_effect=reject_batch_processed_writes + ) + + mock_response = ResponsesAPIResponse( + id="resp_old_schema", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage(input_tokens=100, output_tokens=50, total_tokens=150), + ) + + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: + mock_aget.return_value = mock_response + await check_responses_cost_instance.check_responses_cost() + + mock_aget.assert_awaited_once() + assert _completed_job_ids(mock_prisma_client) == ["job-old-schema"] + + @pytest.mark.asyncio + async def test_terminal_response_replaces_the_queued_file_object_on_the_row( + self, check_responses_cost_instance, mock_prisma_client + ): + """The row has to become the durable copy of the finished generation. Leaving the queued + placeholder there forces the retrieve endpoint back to the provider, and every re-read + replays the same usage under a fresh id, which is what billed the job again per read.""" + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_persisted" + mock_job.created_by = "test-user" + mock_job.id = "job-persisted" + mock_job.file_object = { + "model": "gpt-5", + "id": "resp_test_persisted", + "status": "queued", + "usage": None, + } + + 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_response = ResponsesAPIResponse( + id="resp_finished_upstream", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage( + input_tokens=100, output_tokens=50, total_tokens=150 + ), + ) + + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: + mock_aget.return_value = mock_response + await check_responses_cost_instance.check_responses_cost() + + completion_calls = _completion_calls(mock_prisma_client) + assert len(completion_calls) == 1 + assert completion_calls[0].kwargs["where"] == {"id": "job-persisted"} + + persisted = _persisted_response(completion_calls[0]) + assert persisted["id"] == "resp_finished_upstream" + assert persisted["status"] == "completed" + assert persisted["usage"]["total_tokens"] == 150 + + @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 + ): + """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.""" + mock_job1 = MagicMock() + mock_job1.unified_object_id = "resp_test_persist_fails" + mock_job1.created_by = "user1" + mock_job1.id = "job-persist-fails" + mock_job1.file_object = {"model": "gpt-5", "id": "resp_test_persist_fails"} + + mock_job2 = MagicMock() + mock_job2.unified_object_id = "resp_test_persist_works" + mock_job2.created_by = "user2" + mock_job2.id = "job-persist-works" + mock_job2.file_object = {"model": "gpt-5", "id": "resp_test_persist_works"} + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job1, mock_job2] + ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + side_effect=[1, 1, Exception("deadlock detected"), 1] + ) + + mock_response = ResponsesAPIResponse( + id="resp_persisted", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage(input_tokens=100, output_tokens=50, total_tokens=150), + ) + + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: + mock_aget.return_value = mock_response + await check_responses_cost_instance.check_responses_cost() + + assert mock_aget.await_count == 2 + assert _completed_job_ids(mock_prisma_client) == [ + "job-persist-fails", + "job-persist-works", + ] diff --git a/tests/test_litellm/integrations/test_responses_background_cost.py b/tests/test_litellm/integrations/test_responses_background_cost.py index 7ad39194222..8e5e72bf032 100644 --- a/tests/test_litellm/integrations/test_responses_background_cost.py +++ b/tests/test_litellm/integrations/test_responses_background_cost.py @@ -369,7 +369,9 @@ class TestCheckResponsesCost: ) # Mock update_many - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + return_value=1 + ) # Create a completed response completed_response = ResponsesAPIResponse( @@ -398,17 +400,17 @@ class TestCheckResponsesCost: await checker.check_responses_cost() # Verify update_many was called to mark job as completed - # (stale cleanup also calls update_many, so check the specific completion call) + # (the costing claim also calls update_many, so check the specific completion call) update_many_calls = ( mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list ) completion_calls = [ c for c in update_many_calls - if c.kwargs.get("where", {}).get("id") is not None + if c.kwargs["data"].get("status") == "completed" ] assert len(completion_calls) == 1 - assert completion_calls[0].kwargs["where"]["id"]["in"] == ["job-123"] + assert completion_calls[0].kwargs["where"]["id"] == "job-123" assert completion_calls[0].kwargs["data"]["status"] == "completed" @pytest.mark.asyncio @@ -429,7 +431,9 @@ class TestCheckResponsesCost: mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( return_value=[mock_job] ) - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + return_value=1 + ) # Create a failed response failed_response = ResponsesAPIResponse( @@ -453,14 +457,14 @@ class TestCheckResponsesCost: await checker.check_responses_cost() # Verify job was marked as completed even though it failed - # (stale cleanup also calls update_many, so check the specific completion call) + # (the costing claim also calls update_many, so check the specific completion call) update_many_calls = ( mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list ) completion_calls = [ c for c in update_many_calls - if c.kwargs.get("where", {}).get("id") is not None + if c.kwargs["data"].get("status") == "completed" ] assert len(completion_calls) == 1 @@ -482,7 +486,9 @@ class TestCheckResponsesCost: mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( return_value=[mock_job] ) - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + return_value=1 + ) # Create an in-progress response in_progress_response = ResponsesAPIResponse( @@ -506,14 +512,14 @@ class TestCheckResponsesCost: await checker.check_responses_cost() # Verify no completion update_many was called (job still in progress) - # (stale cleanup may still call update_many, so filter for completion calls) + # (the claim and its release also call update_many, so filter for completion calls) update_many_calls = ( mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list ) completion_calls = [ c for c in update_many_calls - if c.kwargs.get("where", {}).get("id") is not None + if c.kwargs["data"].get("status") == "completed" ] assert len(completion_calls) == 0 @@ -535,7 +541,9 @@ class TestCheckResponsesCost: mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( return_value=[mock_job] ) - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + return_value=1 + ) checker = CheckResponsesCost( proxy_logging_obj=mock_proxy_logging_obj, @@ -553,13 +561,13 @@ class TestCheckResponsesCost: await checker.check_responses_cost() # Verify no completion update_many was called (error occurred) - # (stale cleanup may still call update_many, so filter for completion calls) + # (the claim and its release also call update_many, so filter for completion calls) update_many_calls = ( mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list ) completion_calls = [ c for c in update_many_calls - if c.kwargs.get("where", {}).get("id") is not None + if c.kwargs["data"].get("status") == "completed" ] assert len(completion_calls) == 0 From ffa52f2f6aa9e5b7c5caa69effe191be5b8b3e29 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Wed, 9 Sep 2026 16:02:50 -0700 Subject: [PATCH 03/13] refactor(responses): mark a billed row completed without copying the response onto it The previous commit stored the finished ResponsesAPIResponse in file_object. That duplicates content the provider still serves from its own copy, and the usage and spend it was meant to preserve already land in LiteLLM_SpendLogs on every billed call regardless of store_prompts_in_spend_logs, which gates only the messages and response body columns. The poller now writes status alone, as it did before. The write stays per job rather than one bulk update so a single failure cannot strand the rest of the cycle. Claude-Session: https://claude.ai/code/session_01Hn5E8Jz1LjGLFyiYxBRcBW --- .../common_utils/check_responses_cost.py | 24 ++++----- .../test_check_responses_cost.py | 54 ------------------- 2 files changed, 10 insertions(+), 68 deletions(-) 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 62e742be680..c756832a7b7 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py @@ -178,16 +178,12 @@ class CheckResponsesCost: f"so its cost will not be retried: {db_err}" ) - async def _persist_terminal_response( - self, job: "LiteLLM_ManagedObjectTable", response: ResponsesAPIResponse - ) -> None: - """Store the finished response on its managed row and retire the row from polling. + 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. - The row is the only copy of a background generation's usage that outlives the poll, so - ``GET /v1/responses/{id}`` can serve a terminal job from here instead of re-reading it - from the provider. Every provider re-read replays the same usage and hands back a - freshly encoded id, which is what made this route bill per read and defeated id-based - dedup in the first place. + Only ``status`` is written. The generation's usage and spend already land in + ``LiteLLM_SpendLogs`` unconditionally, so copying the response body onto this row would + duplicate content the provider still serves, on a table nothing ever deletes from. ``status`` stays the literal "completed" for every terminal provider status, matching what this poller has always written, so stale-row expiry keeps skipping these rows. @@ -195,11 +191,11 @@ class CheckResponsesCost: try: await self.prisma_client.db.litellm_managedobjecttable.update_many( where={"id": job.id}, - data={"status": "completed", "file_object": response.model_dump_json()}, + data={"status": "completed"}, ) except Exception as db_err: verbose_proxy_logger.error( - f"CheckResponsesCost: failed to persist terminal response for job {job.id}: {db_err}" + f"CheckResponsesCost: failed to mark job {job.id} completed: {db_err}" ) async def check_responses_cost(self): @@ -293,10 +289,10 @@ class CheckResponsesCost: verbose_proxy_logger.info( f"Response {unified_object_id} has terminal status {response.status}, marking as complete" ) - completed_jobs.append((job, response)) + completed_jobs.append(job) - for job, response in completed_jobs: - await self._persist_terminal_response(job, response) + for job in completed_jobs: + await self._mark_job_completed(job) if len(completed_jobs) > 0: verbose_proxy_logger.info( diff --git a/tests/proxy_unit_tests/test_check_responses_cost.py b/tests/proxy_unit_tests/test_check_responses_cost.py index 2d832975a11..bbf2b88e5b7 100644 --- a/tests/proxy_unit_tests/test_check_responses_cost.py +++ b/tests/proxy_unit_tests/test_check_responses_cost.py @@ -3,7 +3,6 @@ Unit tests for CheckResponsesCost class """ import asyncio -import json from datetime import datetime, timedelta, timezone from unittest.mock import AsyncMock, MagicMock, Mock, patch @@ -31,10 +30,6 @@ def _completed_job_ids(mock_prisma_client): return [call.kwargs["where"]["id"] for call in _completion_calls(mock_prisma_client)] -def _persisted_response(completion_call): - return json.loads(completion_call.kwargs["data"]["file_object"]) - - def _claim_calls(mock_prisma_client): return _update_many_calls_writing( mock_prisma_client, lambda data: data == {"batch_processed": True} @@ -1119,55 +1114,6 @@ class TestCheckResponsesCost: mock_aget.assert_awaited_once() assert _completed_job_ids(mock_prisma_client) == ["job-old-schema"] - @pytest.mark.asyncio - async def test_terminal_response_replaces_the_queued_file_object_on_the_row( - self, check_responses_cost_instance, mock_prisma_client - ): - """The row has to become the durable copy of the finished generation. Leaving the queued - placeholder there forces the retrieve endpoint back to the provider, and every re-read - replays the same usage under a fresh id, which is what billed the job again per read.""" - mock_job = MagicMock() - mock_job.unified_object_id = "resp_test_persisted" - mock_job.created_by = "test-user" - mock_job.id = "job-persisted" - mock_job.file_object = { - "model": "gpt-5", - "id": "resp_test_persisted", - "status": "queued", - "usage": None, - } - - 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_response = ResponsesAPIResponse( - id="resp_finished_upstream", - object="response", - status="completed", - created_at=int(datetime.now().timestamp()), - output=[], - usage=ResponseAPIUsage( - input_tokens=100, output_tokens=50, total_tokens=150 - ), - ) - - with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: - mock_aget.return_value = mock_response - await check_responses_cost_instance.check_responses_cost() - - completion_calls = _completion_calls(mock_prisma_client) - assert len(completion_calls) == 1 - assert completion_calls[0].kwargs["where"] == {"id": "job-persisted"} - - persisted = _persisted_response(completion_calls[0]) - assert persisted["id"] == "resp_finished_upstream" - assert persisted["status"] == "completed" - assert persisted["usage"]["total_tokens"] == 150 - @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 From 49a8ce5e42b20ff8b217198f75aece683da59297 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Wed, 9 Sep 2026 17:57:23 -0700 Subject: [PATCH 04/13] fix(responses): key a background response's managed row by the provider id model_object_id is documented as "the id returned by the backend API provider", and that is what batches and fine-tuning jobs store there. The background responses create stored the advertised id in it instead, which is encrypted with a fresh nonce on every call, so the row had no stable handle on the generation it describes. The cost poller now reads the provider id straight off the row. Rows written before this still carry the advertised id there, and decrypting is a no-op on an id that is already the provider's, so both shapes resolve through the same call. Claude-Session: https://claude.ai/code/session_01RHAjRxNhXTpKHeGMZ1nDKi --- .../common_utils/check_responses_cost.py | 4 +- .../proxy/response_api_endpoints/endpoints.py | 7 +- .../test_check_responses_cost.py | 120 ++++++++++++++++ .../response_api_endpoints/test_endpoints.py | 134 ++++++++++++++++++ 4 files changed, 262 insertions(+), 3 deletions(-) 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 c756832a7b7..47a6f38e8cb 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py @@ -238,8 +238,8 @@ class CheckResponsesCost: stored_response = job.file_object model_name = stored_response.get("model", None) - # Decrypt the response ID - responses_id_security, _, _ = ResponsesIDSecurity()._decrypt_response_id(unified_object_id) + # Decrypts rows written before model_object_id held the provider's own id. + responses_id_security, _, _ = ResponsesIDSecurity()._decrypt_response_id(job.model_object_id) # Prepare metadata with model information for cost tracking litellm_metadata = { diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index f320e2c65fb..95f3615026b 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -379,6 +379,10 @@ async def responses_api( if managed_files_obj and llm_router: try: + from litellm.proxy.hooks.responses_id_security import ( + ResponsesIDSecurity, + ) + # Get the actual deployment model_id from hidden params hidden_params: Final = getattr(response, "_hidden_params", {}) or {} model_id: Final = hidden_params.get("model_id", None) @@ -389,12 +393,13 @@ async def responses_api( response.id, ) raise Exception("No model_id found in response hidden params") + provider_response_id, _, _ = ResponsesIDSecurity()._decrypt_response_id(response.id) # Store in managed objects table await managed_files_obj.store_unified_object_id( unified_object_id=response.id, file_object=response, litellm_parent_otel_span=None, - model_object_id=response.id, + model_object_id=provider_response_id, file_purpose="response", user_api_key_dict=user_api_key_dict, persist_attribution=True, diff --git a/tests/proxy_unit_tests/test_check_responses_cost.py b/tests/proxy_unit_tests/test_check_responses_cost.py index bbf2b88e5b7..351dc6a4cac 100644 --- a/tests/proxy_unit_tests/test_check_responses_cost.py +++ b/tests/proxy_unit_tests/test_check_responses_cost.py @@ -142,6 +142,7 @@ class TestCheckResponsesCost: # Mock job with response ID mock_job = MagicMock() mock_job.unified_object_id = "resp_test_123" + mock_job.model_object_id = "resp_test_123" mock_job.created_by = "test-user" mock_job.id = "job-123" mock_job.file_object = {"model": "gpt-4o", "id": "resp_test_123"} @@ -185,6 +186,7 @@ class TestCheckResponsesCost: # Mock job mock_job = MagicMock() mock_job.unified_object_id = "resp_test_456" + mock_job.model_object_id = "resp_test_456" mock_job.created_by = "test-user" mock_job.id = "job-456" mock_job.file_object = {"model": "gpt-4o", "id": "resp_test_456"} @@ -224,6 +226,7 @@ class TestCheckResponsesCost: # Mock job mock_job = MagicMock() mock_job.unified_object_id = "resp_test_789" + mock_job.model_object_id = "resp_test_789" mock_job.created_by = "test-user" mock_job.id = "job-789" mock_job.file_object = {"model": "gpt-4o", "id": "resp_test_789"} @@ -263,6 +266,7 @@ class TestCheckResponsesCost: # Mock job mock_job = MagicMock() mock_job.unified_object_id = "resp_test_in_progress" + mock_job.model_object_id = "resp_test_in_progress" mock_job.created_by = "test-user" mock_job.id = "job-in-progress" mock_job.file_object = {"model": "gpt-4o", "id": "resp_test_in_progress"} @@ -304,6 +308,7 @@ class TestCheckResponsesCost: # Mock job mock_job = MagicMock() mock_job.unified_object_id = "resp_test_queued" + mock_job.model_object_id = "resp_test_queued" mock_job.created_by = "test-user" mock_job.id = "job-queued" mock_job.file_object = {"model": "gpt-4o", "id": "resp_test_queued"} @@ -345,6 +350,7 @@ class TestCheckResponsesCost: # Mock job mock_job = MagicMock() mock_job.unified_object_id = "resp_test_error" + mock_job.model_object_id = "resp_test_error" mock_job.created_by = "test-user" mock_job.id = "job-error" mock_job.file_object = {"model": "gpt-4o", "id": "resp_test_error"} @@ -379,18 +385,21 @@ class TestCheckResponsesCost: # Mock multiple jobs mock_job1 = MagicMock() mock_job1.unified_object_id = "resp_test_1" + mock_job1.model_object_id = "resp_test_1" mock_job1.created_by = "user1" mock_job1.id = "job-1" mock_job1.file_object = {"model": "gpt-4o", "id": "resp_test_1"} mock_job2 = MagicMock() mock_job2.unified_object_id = "resp_test_2" + mock_job2.model_object_id = "resp_test_2" mock_job2.created_by = "user2" mock_job2.id = "job-2" mock_job2.file_object = {"model": "gpt-4o", "id": "resp_test_2"} mock_job3 = MagicMock() mock_job3.unified_object_id = "resp_test_3" + mock_job3.model_object_id = "resp_test_3" mock_job3.created_by = "user3" mock_job3.id = "job-3" mock_job3.file_object = {"model": "gpt-4o", "id": "resp_test_3"} @@ -470,6 +479,7 @@ class TestCheckResponsesCost: mock_job = MagicMock() mock_job.unified_object_id = encoded_response_id + mock_job.model_object_id = encoded_response_id mock_job.created_by = "test-user" mock_job.id = "job-router" mock_job.file_object = {"model": "azure-gpt-5", "id": encoded_response_id} @@ -541,6 +551,7 @@ class TestCheckResponsesCost: mock_job = MagicMock() mock_job.unified_object_id = encrypted_response_id + mock_job.model_object_id = encrypted_response_id mock_job.created_by = "test-user" mock_job.id = "job-encrypted" mock_job.file_object = {"model": "gpt-5", "id": encrypted_response_id} @@ -586,6 +597,7 @@ class TestCheckResponsesCost: """Ids that carry no deployment info can't be routed, so fall back to the SDK.""" mock_job = MagicMock() mock_job.unified_object_id = "resp_plain_upstream_id" + mock_job.model_object_id = "resp_plain_upstream_id" mock_job.created_by = "test-user" mock_job.id = "job-plain" mock_job.file_object = {"model": "gpt-5", "id": "resp_plain_upstream_id"} @@ -635,6 +647,7 @@ class TestCheckResponsesCost: mock_job = MagicMock() mock_job.unified_object_id = encoded_response_id + mock_job.model_object_id = encoded_response_id mock_job.created_by = "test-user" mock_job.id = "job-missing-deployment" mock_job.file_object = {"model": "gpt-5", "id": encoded_response_id} @@ -677,6 +690,7 @@ class TestCheckResponsesCost: """'incomplete' is terminal in the Responses API, so the row must not stay queued.""" mock_job = MagicMock() mock_job.unified_object_id = "resp_test_incomplete" + mock_job.model_object_id = "resp_test_incomplete" mock_job.created_by = "test-user" mock_job.id = "job-incomplete" mock_job.file_object = {"model": "gpt-5", "id": "resp_test_incomplete"} @@ -710,6 +724,7 @@ class TestCheckResponsesCost: """When file_object has no 'model' key, model_name is None and metadata skips model fields.""" mock_job = MagicMock() mock_job.unified_object_id = "resp_test_no_model" + mock_job.model_object_id = "resp_test_no_model" mock_job.created_by = "test-user" mock_job.team_id = None mock_job.api_key = None @@ -752,6 +767,7 @@ class TestCheckResponsesCost: mock_job = MagicMock() mock_job.unified_object_id = "resp_test_billed" + mock_job.model_object_id = "resp_test_billed" mock_job.created_by = "test-user" mock_job.team_id = "team-billed" mock_job.api_key = "sk-billed" @@ -788,6 +804,7 @@ class TestCheckResponsesCost: spend log, so losing the claim has to skip the read entirely or the job is billed twice.""" mock_job = MagicMock() mock_job.unified_object_id = "resp_test_claimed_elsewhere" + mock_job.model_object_id = "resp_test_claimed_elsewhere" mock_job.created_by = "test-user" mock_job.id = "job-claimed-elsewhere" mock_job.file_object = {"model": "gpt-5", "id": "resp_test_claimed_elsewhere"} @@ -828,6 +845,7 @@ class TestCheckResponsesCost: mock_job = MagicMock() mock_job.unified_object_id = "resp_test_abandoned" + mock_job.model_object_id = "resp_test_abandoned" mock_job.created_by = "test-user" mock_job.id = "job-abandoned" mock_job.file_object = {"model": "gpt-5", "id": "resp_test_abandoned"} @@ -869,6 +887,7 @@ class TestCheckResponsesCost: afterwards is what stops a second pod reading and billing the same row again.""" mock_job = MagicMock() mock_job.unified_object_id = "resp_test_ordering" + mock_job.model_object_id = "resp_test_ordering" mock_job.created_by = "test-user" mock_job.id = "job-ordering" mock_job.file_object = {"model": "gpt-5", "id": "resp_test_ordering"} @@ -919,6 +938,7 @@ class TestCheckResponsesCost: back to batch_processed=False; holding the claim retires it before it is ever billed.""" mock_job = MagicMock() mock_job.unified_object_id = "resp_test_still_running" + mock_job.model_object_id = "resp_test_still_running" mock_job.created_by = "test-user" mock_job.id = "job-still-running" mock_job.file_object = {"model": "gpt-5", "id": "resp_test_still_running"} @@ -959,6 +979,7 @@ class TestCheckResponsesCost: retired unbilled and no later poll cycle ever retries it.""" mock_job = MagicMock() mock_job.unified_object_id = "resp_test_read_error" + mock_job.model_object_id = "resp_test_read_error" mock_job.created_by = "test-user" mock_job.id = "job-read-error" mock_job.file_object = {"model": "gpt-5", "id": "resp_test_read_error"} @@ -993,12 +1014,14 @@ class TestCheckResponsesCost: has to be read and billed in the same cycle.""" mock_job1 = MagicMock() mock_job1.unified_object_id = "resp_test_first" + mock_job1.model_object_id = "resp_test_first" mock_job1.created_by = "user1" mock_job1.id = "job-first" mock_job1.file_object = {"model": "gpt-5", "id": "resp_test_first"} mock_job2 = MagicMock() mock_job2.unified_object_id = "resp_test_second" + mock_job2.model_object_id = "resp_test_second" mock_job2.created_by = "user2" mock_job2.id = "job-second" mock_job2.file_object = {"model": "gpt-5", "id": "resp_test_second"} @@ -1078,6 +1101,7 @@ class TestCheckResponsesCost: (which is what bills it) and the row is still marked completed.""" mock_job = MagicMock() mock_job.unified_object_id = "resp_test_old_schema" + mock_job.model_object_id = "resp_test_old_schema" mock_job.created_by = "test-user" mock_job.id = "job-old-schema" mock_job.file_object = {"model": "gpt-5", "id": "resp_test_old_schema"} @@ -1122,12 +1146,14 @@ class TestCheckResponsesCost: are already read and billed, so losing their write loses their usage for good.""" mock_job1 = MagicMock() mock_job1.unified_object_id = "resp_test_persist_fails" + mock_job1.model_object_id = "resp_test_persist_fails" mock_job1.created_by = "user1" mock_job1.id = "job-persist-fails" mock_job1.file_object = {"model": "gpt-5", "id": "resp_test_persist_fails"} mock_job2 = MagicMock() mock_job2.unified_object_id = "resp_test_persist_works" + mock_job2.model_object_id = "resp_test_persist_works" mock_job2.created_by = "user2" mock_job2.id = "job-persist-works" mock_job2.file_object = {"model": "gpt-5", "id": "resp_test_persist_works"} @@ -1157,3 +1183,97 @@ class TestCheckResponsesCost: "job-persist-fails", "job-persist-works", ] + + @pytest.mark.asyncio + async def test_poller_fetches_the_provider_id_from_model_object_id( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router, monkeypatch + ): + """The row's provider id drives the fetch, not the nonce-encrypted advertised id. + + A background create advertises a freshly encrypted id per call, so unified_object_id + is not a stable handle on the generation. + """ + from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper + from litellm.types.utils import SpecialEnums + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-key-for-response-ids") + + provider_response_id = "resp_provider_stable_1" + stale_advertised_id = "resp_" + str( + encrypt_value_helper( + value=SpecialEnums.LITELLM_MANAGED_RESPONSE_API_RESPONSE_ID_COMPLETE_STR.value.format( + "resp_some_other_encoding", "test-user", "test-team" + ) + ) + ) + + mock_job = MagicMock() + mock_job.unified_object_id = stale_advertised_id + mock_job.model_object_id = provider_response_id + mock_job.created_by = "test-user" + mock_job.id = "job-provider-id" + mock_job.file_object = {"model": "gpt-5", "id": stale_advertised_id} + + 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_response = ResponsesAPIResponse( + id=provider_response_id, + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage(input_tokens=10, output_tokens=5, total_tokens=15), + ) + + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: + mock_aget.return_value = mock_response + await check_responses_cost_instance.check_responses_cost() + + assert mock_aget.call_args[1]["response_id"] == provider_response_id + assert _completed_job_ids(mock_prisma_client) == ["job-provider-id"] + + @pytest.mark.asyncio + async def test_poller_still_reads_rows_written_before_the_provider_id_was_stored( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router, monkeypatch + ): + """Rows created earlier carry the encrypted advertised id in both columns.""" + from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper + from litellm.types.utils import SpecialEnums + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-key-for-response-ids") + + provider_response_id = "resp_legacy_upstream_9" + legacy_id = "resp_" + str( + encrypt_value_helper( + value=SpecialEnums.LITELLM_MANAGED_RESPONSE_API_RESPONSE_ID_COMPLETE_STR.value.format( + provider_response_id, "test-user", "test-team" + ) + ) + ) + + mock_job = MagicMock() + mock_job.unified_object_id = legacy_id + mock_job.model_object_id = legacy_id + mock_job.created_by = "test-user" + mock_job.id = "job-legacy" + mock_job.file_object = {"model": "gpt-5", "id": legacy_id} + + 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_response = ResponsesAPIResponse( + id=provider_response_id, + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage(input_tokens=10, output_tokens=5, total_tokens=15), + ) + + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: + mock_aget.return_value = mock_response + await check_responses_cost_instance.check_responses_cost() + + assert mock_aget.call_args[1]["response_id"] == provider_response_id + assert _completed_job_ids(mock_prisma_client) == ["job-legacy"] diff --git a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py index d7010de6405..f120f88bba6 100644 --- a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py @@ -1959,3 +1959,137 @@ class TestResponsesInputTokens: assert response.status_code == 429, response.text assert response.json()["error"]["message"] == "rate limited" + + +class TestBackgroundResponseManagedObjectId: + """The managed row for a background response must be keyed by the provider's own id. + + The advertised ``response.id`` is encrypted with a fresh nonce per call, so storing it + in ``model_object_id`` leaves the row with no stable lookup key and every later read + of the same generation looks like a new object. + """ + + @staticmethod + def _encrypted_id(provider_response_id: str) -> str: + from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper + from litellm.types.utils import SpecialEnums + + managed_id = SpecialEnums.LITELLM_MANAGED_RESPONSE_API_RESPONSE_ID_COMPLETE_STR.value.format( + provider_response_id, "u-1", "t-1" + ) + return f"resp_{encrypt_value_helper(value=managed_id)}" + + async def _store_call_for(self, provider_response_id: str) -> dict: + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.response_api_endpoints.endpoints import responses_api + from litellm.types.llms.openai import ResponsesAPIResponse + + advertised_id = self._encrypted_id(provider_response_id) + assert advertised_id != self._encrypted_id(provider_response_id), ( + "advertised ids must be nonce-encrypted, otherwise this regression cannot occur" + ) + + response = ResponsesAPIResponse( + id=advertised_id, + created_at=0, + model="gpt-4o", + object="response", + output=[], + parallel_tool_calls=False, + tool_choice="auto", + tools=[], + status="queued", + ) + response._hidden_params = {"model_id": "deployment-1"} + + managed_files_obj = MagicMock() + managed_files_obj.store_unified_object_id = AsyncMock() + proxy_logging_obj = MagicMock() + proxy_logging_obj.get_proxy_hook = MagicMock(return_value=managed_files_obj) + + with patch( + "litellm.proxy.proxy_server._read_request_body", + AsyncMock(return_value={"model": "gpt-4o", "input": "hi", "background": True}), + ), patch("litellm.proxy.proxy_server.polling_via_cache_enabled", False), patch( + "litellm.proxy.proxy_server.llm_router", MagicMock() + ), patch( + "litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj + ), patch( + "litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing.base_process_llm_request", + AsyncMock(return_value=response), + ): + await responses_api( + request=MagicMock(), + fastapi_response=MagicMock(), + user_api_key_dict=UserAPIKeyAuth(api_key="sk-1234", user_id="u-1", team_id="t-1"), + ) + + managed_files_obj.store_unified_object_id.assert_awaited_once() + return managed_files_obj.store_unified_object_id.await_args.kwargs + + @pytest.mark.asyncio + async def test_model_object_id_is_the_provider_response_id(self, monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-regression-salt") + provider_response_id = "resp_provider68abc123" + + kwargs = await self._store_call_for(provider_response_id) + + assert kwargs["model_object_id"] == provider_response_id + assert kwargs["unified_object_id"] != provider_response_id + assert kwargs["unified_object_id"] == kwargs["file_object"].id + + @pytest.mark.asyncio + async def test_two_background_creates_are_distinguishable_by_provider_id(self, monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-regression-salt") + + first = await self._store_call_for("resp_providerAAA") + second = await self._store_call_for("resp_providerBBB") + + assert first["model_object_id"] == "resp_providerAAA" + assert second["model_object_id"] == "resp_providerBBB" + + @pytest.mark.asyncio + async def test_unencrypted_advertised_id_is_stored_as_is(self, monkeypatch): + """With response-id security disabled the advertised id is already the provider's.""" + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.response_api_endpoints.endpoints import responses_api + from litellm.types.llms.openai import ResponsesAPIResponse + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-regression-salt") + response = ResponsesAPIResponse( + id="resp_rawprovider999", + created_at=0, + model="gpt-4o", + object="response", + output=[], + parallel_tool_calls=False, + tool_choice="auto", + tools=[], + status="queued", + ) + response._hidden_params = {"model_id": "deployment-1"} + + managed_files_obj = MagicMock() + managed_files_obj.store_unified_object_id = AsyncMock() + proxy_logging_obj = MagicMock() + proxy_logging_obj.get_proxy_hook = MagicMock(return_value=managed_files_obj) + + with patch( + "litellm.proxy.proxy_server._read_request_body", + AsyncMock(return_value={"model": "gpt-4o", "input": "hi", "background": True}), + ), patch("litellm.proxy.proxy_server.polling_via_cache_enabled", False), patch( + "litellm.proxy.proxy_server.llm_router", MagicMock() + ), patch( + "litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj + ), patch( + "litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing.base_process_llm_request", + AsyncMock(return_value=response), + ): + await responses_api( + request=MagicMock(), + fastapi_response=MagicMock(), + user_api_key_dict=UserAPIKeyAuth(api_key="sk-1234", user_id="u-1", team_id="t-1"), + ) + + kwargs = managed_files_obj.store_unified_object_id.await_args.kwargs + assert kwargs["model_object_id"] == "resp_rawprovider999" From 492336a50b0781fd5f8ef0b4df448834703ccdc3 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Wed, 9 Sep 2026 18:25:07 -0700 Subject: [PATCH 05/13] refactor(responses): give the background row store an injectable seam The store lived inline in responses_api, so covering it meant patching five proxy_server globals per test, which the test-quality gate counts as pinning the test to the wiring rather than the behaviour. It is now a module-level function that takes the managed-files hook as an argument, and the tests hand it a fake directly. The poller tests reach the same fetch through the router fixture that is already injected, instead of patching litellm.aget_responses. A response with no model_id now returns early rather than raising into the caller's except block. Nothing is stored either way and the warning is unchanged, so only the redundant second log line goes away. Claude-Session: https://claude.ai/code/session_01RHAjRxNhXTpKHeGMZ1nDKi --- .../proxy/response_api_endpoints/endpoints.py | 104 +++++++------- .../test_check_responses_cost.py | 64 +++++---- .../response_api_endpoints/test_endpoints.py | 130 +++++++----------- 3 files changed, 147 insertions(+), 151 deletions(-) diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index 95f3615026b..9f0e4937bba 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -39,10 +39,47 @@ from litellm.types.responses.main import DeleteResponseResult from litellm.types.utils import TokenCountResponse if TYPE_CHECKING: + from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles + from litellm.router import Router router: Final = APIRouter() + +async def store_background_response_object( + response: ResponsesAPIResponse, + managed_files_obj: "_PROXY_LiteLLMManagedFiles", + user_api_key_dict: UserAPIKeyAuth, +) -> None: + """Record a queued background response so the cost poller can find and bill it. + + ``model_object_id`` carries the provider's own id because the advertised ``response.id`` + is re-encrypted with a fresh nonce on every call, leaving the row no stable handle on + the generation it describes. + """ + from litellm.proxy.hooks.responses_id_security import ResponsesIDSecurity + + hidden_params: Final = getattr(response, "_hidden_params", {}) or {} + if not hidden_params.get("model_id"): + verbose_proxy_logger.warning( + "No model_id found in response hidden params for response %s, skipping managed object storage", + response.id, + ) + return + + provider_response_id, _, _ = ResponsesIDSecurity()._decrypt_response_id(response.id) + await managed_files_obj.store_unified_object_id( + unified_object_id=response.id, + file_object=response, + litellm_parent_otel_span=None, + model_object_id=provider_response_id, + file_purpose="response", + user_api_key_dict=user_api_key_dict, + persist_attribution=True, + ) + verbose_proxy_logger.info("Stored background response %s in managed objects table", response.id) + + _user_api_key_auth_dep: Final = Depends(user_api_key_auth) _RESPONSES_TAGS: Final[list[str | Enum]] = ["responses"] # mutable-ok: fastapi's route signature requires list tags @@ -366,54 +403,29 @@ async def responses_api( ) # Store in managed objects table if background mode is enabled - if data.get("background") and isinstance(response, ResponsesAPIResponse): - if response.status in ["queued", "in_progress"]: - from litellm_enterprise.proxy.hooks.managed_files import ( - _PROXY_LiteLLMManagedFiles, - ) + if ( + data.get("background") + and isinstance(response, ResponsesAPIResponse) + and response.status in ("queued", "in_progress") + ): + from litellm_enterprise.proxy.hooks.managed_files import ( + _PROXY_LiteLLMManagedFiles, + ) - managed_files_obj: Final = cast( - _PROXY_LiteLLMManagedFiles | None, - proxy_logging_obj.get_proxy_hook("managed_files"), - ) + managed_files_obj: Final = cast( + _PROXY_LiteLLMManagedFiles | None, + proxy_logging_obj.get_proxy_hook("managed_files"), + ) - if managed_files_obj and llm_router: - try: - from litellm.proxy.hooks.responses_id_security import ( - ResponsesIDSecurity, - ) - - # Get the actual deployment model_id from hidden params - hidden_params: Final = getattr(response, "_hidden_params", {}) or {} - model_id: Final = hidden_params.get("model_id", None) - - if not model_id: - verbose_proxy_logger.warning( - "No model_id found in response hidden params for response %s, skipping managed object storage", - response.id, - ) - raise Exception("No model_id found in response hidden params") - provider_response_id, _, _ = ResponsesIDSecurity()._decrypt_response_id(response.id) - # Store in managed objects table - await managed_files_obj.store_unified_object_id( - unified_object_id=response.id, - file_object=response, - litellm_parent_otel_span=None, - model_object_id=provider_response_id, - file_purpose="response", - user_api_key_dict=user_api_key_dict, - persist_attribution=True, - ) - - verbose_proxy_logger.info( - "Stored background response %s in managed objects table with unified_id=%s", - response.id, - response.id, - ) - except Exception as e: - verbose_proxy_logger.error( - "Failed to store background response in managed objects table: %s", e - ) + if managed_files_obj and llm_router: + try: + await store_background_response_object( + response=response, + managed_files_obj=managed_files_obj, + user_api_key_dict=user_api_key_dict, + ) + except Exception as e: + verbose_proxy_logger.error("Failed to store background response in managed objects table: %s", e) return response except ModifyResponseException as e: diff --git a/tests/proxy_unit_tests/test_check_responses_cost.py b/tests/proxy_unit_tests/test_check_responses_cost.py index 351dc6a4cac..4768cdc576d 100644 --- a/tests/proxy_unit_tests/test_check_responses_cost.py +++ b/tests/proxy_unit_tests/test_check_responses_cost.py @@ -1191,18 +1191,23 @@ class TestCheckResponsesCost: """The row's provider id drives the fetch, not the nonce-encrypted advertised id. A background create advertises a freshly encrypted id per call, so unified_object_id - is not a stable handle on the generation. + is no handle on the generation. """ from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper + from litellm.responses.utils import ResponsesAPIRequestUtils from litellm.types.utils import SpecialEnums monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-key-for-response-ids") - provider_response_id = "resp_provider_stable_1" + provider_response_id = ResponsesAPIRequestUtils._build_responses_api_response_id( + custom_llm_provider="openai", + model_id="deployment-xyz", + response_id="resp_upstream_stable", + ) stale_advertised_id = "resp_" + str( encrypt_value_helper( value=SpecialEnums.LITELLM_MANAGED_RESPONSE_API_RESPONSE_ID_COMPLETE_STR.value.format( - "resp_some_other_encoding", "test-user", "test-team" + "resp_a_previous_encoding", "test-user", "test-team" ) ) ) @@ -1216,21 +1221,20 @@ class TestCheckResponsesCost: 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_response = ResponsesAPIResponse( - id=provider_response_id, - object="response", - status="completed", - created_at=int(datetime.now().timestamp()), - output=[], - usage=ResponseAPIUsage(input_tokens=10, output_tokens=5, total_tokens=15), + mock_llm_router.aget_responses = AsyncMock( + return_value=ResponsesAPIResponse( + id=provider_response_id, + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage(input_tokens=10, output_tokens=5, total_tokens=15), + ) ) - with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: - mock_aget.return_value = mock_response - await check_responses_cost_instance.check_responses_cost() + await check_responses_cost_instance.check_responses_cost() - assert mock_aget.call_args[1]["response_id"] == provider_response_id + assert mock_llm_router.aget_responses.call_args[1]["response_id"] == provider_response_id assert _completed_job_ids(mock_prisma_client) == ["job-provider-id"] @pytest.mark.asyncio @@ -1239,11 +1243,16 @@ class TestCheckResponsesCost: ): """Rows created earlier carry the encrypted advertised id in both columns.""" from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper + from litellm.responses.utils import ResponsesAPIRequestUtils from litellm.types.utils import SpecialEnums monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-key-for-response-ids") - provider_response_id = "resp_legacy_upstream_9" + provider_response_id = ResponsesAPIRequestUtils._build_responses_api_response_id( + custom_llm_provider="openai", + model_id="deployment-xyz", + response_id="resp_legacy_upstream", + ) legacy_id = "resp_" + str( encrypt_value_helper( value=SpecialEnums.LITELLM_MANAGED_RESPONSE_API_RESPONSE_ID_COMPLETE_STR.value.format( @@ -1261,19 +1270,18 @@ class TestCheckResponsesCost: 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_response = ResponsesAPIResponse( - id=provider_response_id, - object="response", - status="completed", - created_at=int(datetime.now().timestamp()), - output=[], - usage=ResponseAPIUsage(input_tokens=10, output_tokens=5, total_tokens=15), + mock_llm_router.aget_responses = AsyncMock( + return_value=ResponsesAPIResponse( + id=provider_response_id, + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage(input_tokens=10, output_tokens=5, total_tokens=15), + ) ) - with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: - mock_aget.return_value = mock_response - await check_responses_cost_instance.check_responses_cost() + await check_responses_cost_instance.check_responses_cost() - assert mock_aget.call_args[1]["response_id"] == provider_response_id + assert mock_llm_router.aget_responses.call_args[1]["response_id"] == provider_response_id assert _completed_job_ids(mock_prisma_client) == ["job-legacy"] diff --git a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py index f120f88bba6..f528dc6ba58 100644 --- a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py @@ -1962,11 +1962,11 @@ class TestResponsesInputTokens: class TestBackgroundResponseManagedObjectId: - """The managed row for a background response must be keyed by the provider's own id. + """The managed row for a background response is keyed by the provider's own id. The advertised ``response.id`` is encrypted with a fresh nonce per call, so storing it - in ``model_object_id`` leaves the row with no stable lookup key and every later read - of the same generation looks like a new object. + in ``model_object_id`` leaves the row with no stable handle on the generation and every + later read of the same generation looks like a new object. """ @staticmethod @@ -1979,16 +1979,10 @@ class TestBackgroundResponseManagedObjectId: ) return f"resp_{encrypt_value_helper(value=managed_id)}" - async def _store_call_for(self, provider_response_id: str) -> dict: - from litellm.proxy._types import UserAPIKeyAuth - from litellm.proxy.response_api_endpoints.endpoints import responses_api + @staticmethod + def _queued_response(advertised_id: str, model_id: str | None = "deployment-1"): from litellm.types.llms.openai import ResponsesAPIResponse - advertised_id = self._encrypted_id(provider_response_id) - assert advertised_id != self._encrypted_id(provider_response_id), ( - "advertised ids must be nonce-encrypted, otherwise this regression cannot occur" - ) - response = ResponsesAPIResponse( id=advertised_id, created_at=0, @@ -2000,50 +1994,60 @@ class TestBackgroundResponseManagedObjectId: tools=[], status="queued", ) - response._hidden_params = {"model_id": "deployment-1"} + response._hidden_params = {"model_id": model_id} if model_id else {} + return response + + async def _stored_kwargs(self, advertised_id: str, model_id: str | None = "deployment-1"): + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.response_api_endpoints.endpoints import ( + store_background_response_object, + ) managed_files_obj = MagicMock() managed_files_obj.store_unified_object_id = AsyncMock() - proxy_logging_obj = MagicMock() - proxy_logging_obj.get_proxy_hook = MagicMock(return_value=managed_files_obj) - with patch( - "litellm.proxy.proxy_server._read_request_body", - AsyncMock(return_value={"model": "gpt-4o", "input": "hi", "background": True}), - ), patch("litellm.proxy.proxy_server.polling_via_cache_enabled", False), patch( - "litellm.proxy.proxy_server.llm_router", MagicMock() - ), patch( - "litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj - ), patch( - "litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing.base_process_llm_request", - AsyncMock(return_value=response), - ): - await responses_api( - request=MagicMock(), - fastapi_response=MagicMock(), - user_api_key_dict=UserAPIKeyAuth(api_key="sk-1234", user_id="u-1", team_id="t-1"), - ) - - managed_files_obj.store_unified_object_id.assert_awaited_once() - return managed_files_obj.store_unified_object_id.await_args.kwargs + await store_background_response_object( + response=self._queued_response(advertised_id, model_id), + managed_files_obj=managed_files_obj, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-1234", user_id="u-1", team_id="t-1"), + ) + return managed_files_obj.store_unified_object_id @pytest.mark.asyncio async def test_model_object_id_is_the_provider_response_id(self, monkeypatch): monkeypatch.setenv("LITELLM_SALT_KEY", "sk-regression-salt") provider_response_id = "resp_provider68abc123" + advertised_id = self._encrypted_id(provider_response_id) + assert advertised_id != self._encrypted_id(provider_response_id), ( + "advertised ids must be nonce-encrypted, otherwise this regression cannot occur" + ) - kwargs = await self._store_call_for(provider_response_id) + store = await self._stored_kwargs(advertised_id) + store.assert_awaited_once() + kwargs = store.await_args.kwargs assert kwargs["model_object_id"] == provider_response_id - assert kwargs["unified_object_id"] != provider_response_id - assert kwargs["unified_object_id"] == kwargs["file_object"].id + assert kwargs["unified_object_id"] == advertised_id + assert kwargs["file_object"].id == advertised_id @pytest.mark.asyncio - async def test_two_background_creates_are_distinguishable_by_provider_id(self, monkeypatch): + async def test_two_creates_of_one_generation_share_a_provider_id(self, monkeypatch): + """Re-encrypting the same generation must not look like a second object.""" + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-regression-salt") + provider_response_id = "resp_provider_same_gen" + + first = (await self._stored_kwargs(self._encrypted_id(provider_response_id))).await_args.kwargs + second = (await self._stored_kwargs(self._encrypted_id(provider_response_id))).await_args.kwargs + + assert first["unified_object_id"] != second["unified_object_id"] + assert first["model_object_id"] == second["model_object_id"] == provider_response_id + + @pytest.mark.asyncio + async def test_distinct_generations_keep_distinct_provider_ids(self, monkeypatch): monkeypatch.setenv("LITELLM_SALT_KEY", "sk-regression-salt") - first = await self._store_call_for("resp_providerAAA") - second = await self._store_call_for("resp_providerBBB") + first = (await self._stored_kwargs(self._encrypted_id("resp_providerAAA"))).await_args.kwargs + second = (await self._stored_kwargs(self._encrypted_id("resp_providerBBB"))).await_args.kwargs assert first["model_object_id"] == "resp_providerAAA" assert second["model_object_id"] == "resp_providerBBB" @@ -2051,45 +2055,17 @@ class TestBackgroundResponseManagedObjectId: @pytest.mark.asyncio async def test_unencrypted_advertised_id_is_stored_as_is(self, monkeypatch): """With response-id security disabled the advertised id is already the provider's.""" - from litellm.proxy._types import UserAPIKeyAuth - from litellm.proxy.response_api_endpoints.endpoints import responses_api - from litellm.types.llms.openai import ResponsesAPIResponse - monkeypatch.setenv("LITELLM_SALT_KEY", "sk-regression-salt") - response = ResponsesAPIResponse( - id="resp_rawprovider999", - created_at=0, - model="gpt-4o", - object="response", - output=[], - parallel_tool_calls=False, - tool_choice="auto", - tools=[], - status="queued", - ) - response._hidden_params = {"model_id": "deployment-1"} - managed_files_obj = MagicMock() - managed_files_obj.store_unified_object_id = AsyncMock() - proxy_logging_obj = MagicMock() - proxy_logging_obj.get_proxy_hook = MagicMock(return_value=managed_files_obj) + store = await self._stored_kwargs("resp_rawprovider999") - with patch( - "litellm.proxy.proxy_server._read_request_body", - AsyncMock(return_value={"model": "gpt-4o", "input": "hi", "background": True}), - ), patch("litellm.proxy.proxy_server.polling_via_cache_enabled", False), patch( - "litellm.proxy.proxy_server.llm_router", MagicMock() - ), patch( - "litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj - ), patch( - "litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing.base_process_llm_request", - AsyncMock(return_value=response), - ): - await responses_api( - request=MagicMock(), - fastapi_response=MagicMock(), - user_api_key_dict=UserAPIKeyAuth(api_key="sk-1234", user_id="u-1", team_id="t-1"), - ) + assert store.await_args.kwargs["model_object_id"] == "resp_rawprovider999" - kwargs = managed_files_obj.store_unified_object_id.await_args.kwargs - assert kwargs["model_object_id"] == "resp_rawprovider999" + @pytest.mark.asyncio + async def test_response_without_a_deployment_is_not_stored(self, monkeypatch): + """No model_id means the poller could never route the read, so no row is written.""" + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-regression-salt") + + store = await self._stored_kwargs(self._encrypted_id("resp_no_deployment"), model_id=None) + + store.assert_not_awaited() From 5aa19012495a2bca0fb9d768fedfb38149ff7517 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Wed, 9 Sep 2026 18:31:10 -0700 Subject: [PATCH 06/13] test(responses): drive the poller tests through the injected router The eight claim/release tests reached into `litellm.aget_responses` with `patch`, which the test-quality gate flags and which couples them to an import path rather than to the poller's own seam. Each job now carries a LiteLLM-encoded provider id, so the read routes through the router the fixture already injects and the assertions run against that mock. Claude-Session: https://claude.ai/code/session_01RHAjRxNhXTpKHeGMZ1nDKi --- .../test_check_responses_cost.py | 119 +++++++++--------- 1 file changed, 58 insertions(+), 61 deletions(-) diff --git a/tests/proxy_unit_tests/test_check_responses_cost.py b/tests/proxy_unit_tests/test_check_responses_cost.py index 4768cdc576d..1e7ca7dd85d 100644 --- a/tests/proxy_unit_tests/test_check_responses_cost.py +++ b/tests/proxy_unit_tests/test_check_responses_cost.py @@ -42,6 +42,17 @@ def _release_calls(mock_prisma_client): ) +def _routed_response_id(provider_response_id): + """A LiteLLM-encoded id names a deployment, which is what sends the poll's read through the router.""" + from litellm.responses.utils import ResponsesAPIRequestUtils + + return ResponsesAPIRequestUtils._build_responses_api_response_id( + custom_llm_provider="openai", + model_id="deployment-xyz", + response_id=provider_response_id, + ) + + class TestCheckResponsesCost: """Test suite for CheckResponsesCost class""" @@ -804,7 +815,7 @@ class TestCheckResponsesCost: spend log, so losing the claim has to skip the read entirely or the job is billed twice.""" mock_job = MagicMock() mock_job.unified_object_id = "resp_test_claimed_elsewhere" - mock_job.model_object_id = "resp_test_claimed_elsewhere" + mock_job.model_object_id = _routed_response_id("resp_test_claimed_elsewhere") mock_job.created_by = "test-user" mock_job.id = "job-claimed-elsewhere" mock_job.file_object = {"model": "gpt-5", "id": "resp_test_claimed_elsewhere"} @@ -816,10 +827,8 @@ class TestCheckResponsesCost: return_value=0 ) - with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_sdk_aget: - await check_responses_cost_instance.check_responses_cost() + await check_responses_cost_instance.check_responses_cost() - mock_sdk_aget.assert_not_awaited() mock_llm_router.aget_responses.assert_not_awaited() assert _completion_calls(mock_prisma_client) == [] assert _release_calls(mock_prisma_client) == [] @@ -832,7 +841,7 @@ class TestCheckResponsesCost: @pytest.mark.asyncio async def test_claim_is_taken_back_from_a_pod_that_died_holding_it( - self, check_responses_cost_instance, mock_prisma_client + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router ): """A pod that dies between claiming and billing releases nothing, and the row's status never reaches terminal, so without a lease every later cycle re-selects it and loses. @@ -845,7 +854,7 @@ class TestCheckResponsesCost: mock_job = MagicMock() mock_job.unified_object_id = "resp_test_abandoned" - mock_job.model_object_id = "resp_test_abandoned" + mock_job.model_object_id = _routed_response_id("resp_test_abandoned") mock_job.created_by = "test-user" mock_job.id = "job-abandoned" mock_job.file_object = {"model": "gpt-5", "id": "resp_test_abandoned"} @@ -866,9 +875,9 @@ class TestCheckResponsesCost: usage=ResponseAPIUsage(input_tokens=100, output_tokens=50, total_tokens=150), ) - with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: - mock_aget.return_value = mock_response - await check_responses_cost_instance.check_responses_cost() + mock_llm_router.aget_responses = AsyncMock(return_value=mock_response) + + await check_responses_cost_instance.check_responses_cost() claim_where = _claim_calls(mock_prisma_client)[0].kwargs["where"] abandoned_arm = next(arm for arm in claim_where["OR"] if "updated_at" in arm) @@ -881,13 +890,13 @@ class TestCheckResponsesCost: @pytest.mark.asyncio async def test_claim_is_taken_before_the_billing_read_and_kept_on_a_terminal_status( - self, check_responses_cost_instance, mock_prisma_client + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router ): """The read prices the job, so the claim has to be taken before it, and keeping the claim afterwards is what stops a second pod reading and billing the same row again.""" mock_job = MagicMock() mock_job.unified_object_id = "resp_test_ordering" - mock_job.model_object_id = "resp_test_ordering" + mock_job.model_object_id = _routed_response_id("resp_test_ordering") mock_job.created_by = "test-user" mock_job.id = "job-ordering" mock_job.file_object = {"model": "gpt-5", "id": "resp_test_ordering"} @@ -919,10 +928,9 @@ class TestCheckResponsesCost: side_effect=record_update_many ) - with patch( - "litellm.aget_responses", new_callable=AsyncMock, side_effect=record_read - ): - await check_responses_cost_instance.check_responses_cost() + mock_llm_router.aget_responses = AsyncMock(side_effect=record_read) + + await check_responses_cost_instance.check_responses_cost() assert len(writes_and_reads) == 3 assert writes_and_reads[0] == {"batch_processed": True} @@ -932,13 +940,13 @@ class TestCheckResponsesCost: @pytest.mark.asyncio @pytest.mark.parametrize("provider_status", ["queued", "in_progress"]) async def test_non_terminal_status_releases_the_claim( - self, check_responses_cost_instance, mock_prisma_client, provider_status + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router, provider_status ): """A response the provider has not finished yet has no spend to record, so its row must go back to batch_processed=False; holding the claim retires it before it is ever billed.""" mock_job = MagicMock() mock_job.unified_object_id = "resp_test_still_running" - mock_job.model_object_id = "resp_test_still_running" + mock_job.model_object_id = _routed_response_id("resp_test_still_running") mock_job.created_by = "test-user" mock_job.id = "job-still-running" mock_job.file_object = {"model": "gpt-5", "id": "resp_test_still_running"} @@ -959,9 +967,9 @@ class TestCheckResponsesCost: usage=None, ) - with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: - mock_aget.return_value = mock_response - await check_responses_cost_instance.check_responses_cost() + mock_llm_router.aget_responses = AsyncMock(return_value=mock_response) + + await check_responses_cost_instance.check_responses_cost() assert _completion_calls(mock_prisma_client) == [] release_calls = _release_calls(mock_prisma_client) @@ -973,13 +981,13 @@ class TestCheckResponsesCost: @pytest.mark.asyncio async def test_failed_provider_read_releases_the_claim( - self, check_responses_cost_instance, mock_prisma_client + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router ): """A read that raised billed nothing, so the claim has to be handed back or the row is retired unbilled and no later poll cycle ever retries it.""" mock_job = MagicMock() mock_job.unified_object_id = "resp_test_read_error" - mock_job.model_object_id = "resp_test_read_error" + mock_job.model_object_id = _routed_response_id("resp_test_read_error") mock_job.created_by = "test-user" mock_job.id = "job-read-error" mock_job.file_object = {"model": "gpt-5", "id": "resp_test_read_error"} @@ -991,12 +999,9 @@ class TestCheckResponsesCost: return_value=1 ) - with patch( - "litellm.aget_responses", - new_callable=AsyncMock, - side_effect=Exception("Provider error"), - ): - await check_responses_cost_instance.check_responses_cost() + mock_llm_router.aget_responses = AsyncMock(side_effect=Exception("Provider error")) + + await check_responses_cost_instance.check_responses_cost() assert _completion_calls(mock_prisma_client) == [] release_calls = _release_calls(mock_prisma_client) @@ -1008,20 +1013,20 @@ class TestCheckResponsesCost: @pytest.mark.asyncio async def test_a_job_claimed_elsewhere_does_not_block_the_next_job( - self, check_responses_cost_instance, mock_prisma_client + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router ): """Losing one row to another pod must skip only that row: the rest of the poll page still has to be read and billed in the same cycle.""" mock_job1 = MagicMock() mock_job1.unified_object_id = "resp_test_first" - mock_job1.model_object_id = "resp_test_first" + mock_job1.model_object_id = _routed_response_id("resp_test_first") mock_job1.created_by = "user1" mock_job1.id = "job-first" mock_job1.file_object = {"model": "gpt-5", "id": "resp_test_first"} mock_job2 = MagicMock() mock_job2.unified_object_id = "resp_test_second" - mock_job2.model_object_id = "resp_test_second" + mock_job2.model_object_id = _routed_response_id("resp_test_second") mock_job2.created_by = "user2" mock_job2.id = "job-second" mock_job2.file_object = {"model": "gpt-5", "id": "resp_test_second"} @@ -1042,12 +1047,14 @@ class TestCheckResponsesCost: usage=ResponseAPIUsage(input_tokens=100, output_tokens=50, total_tokens=150), ) - with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: - mock_aget.return_value = mock_response - await check_responses_cost_instance.check_responses_cost() + mock_llm_router.aget_responses = AsyncMock(return_value=mock_response) - mock_aget.assert_awaited_once() - assert mock_aget.await_args.kwargs["response_id"] == "resp_test_second" + await check_responses_cost_instance.check_responses_cost() + + mock_llm_router.aget_responses.assert_awaited_once() + assert mock_llm_router.aget_responses.await_args.kwargs["response_id"] == _routed_response_id( + "resp_test_second" + ) assert _completed_job_ids(mock_prisma_client) == ["job-second"] @@ -1095,13 +1102,13 @@ class TestCheckResponsesCost: @pytest.mark.asyncio async def test_old_schema_without_the_claim_column_still_bills_and_completes( - self, check_responses_cost_instance, mock_prisma_client + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router ): """End to end on a pre-migration schema: the claim write fails, the response is still read (which is what bills it) and the row is still marked completed.""" mock_job = MagicMock() mock_job.unified_object_id = "resp_test_old_schema" - mock_job.model_object_id = "resp_test_old_schema" + mock_job.model_object_id = _routed_response_id("resp_test_old_schema") mock_job.created_by = "test-user" mock_job.id = "job-old-schema" mock_job.file_object = {"model": "gpt-5", "id": "resp_test_old_schema"} @@ -1131,29 +1138,29 @@ class TestCheckResponsesCost: usage=ResponseAPIUsage(input_tokens=100, output_tokens=50, total_tokens=150), ) - with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: - mock_aget.return_value = mock_response - await check_responses_cost_instance.check_responses_cost() + mock_llm_router.aget_responses = AsyncMock(return_value=mock_response) - mock_aget.assert_awaited_once() + await check_responses_cost_instance.check_responses_cost() + + mock_llm_router.aget_responses.assert_awaited_once() 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 + 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.""" mock_job1 = MagicMock() mock_job1.unified_object_id = "resp_test_persist_fails" - mock_job1.model_object_id = "resp_test_persist_fails" + mock_job1.model_object_id = _routed_response_id("resp_test_persist_fails") mock_job1.created_by = "user1" mock_job1.id = "job-persist-fails" mock_job1.file_object = {"model": "gpt-5", "id": "resp_test_persist_fails"} mock_job2 = MagicMock() mock_job2.unified_object_id = "resp_test_persist_works" - mock_job2.model_object_id = "resp_test_persist_works" + mock_job2.model_object_id = _routed_response_id("resp_test_persist_works") mock_job2.created_by = "user2" mock_job2.id = "job-persist-works" mock_job2.file_object = {"model": "gpt-5", "id": "resp_test_persist_works"} @@ -1174,11 +1181,11 @@ class TestCheckResponsesCost: usage=ResponseAPIUsage(input_tokens=100, output_tokens=50, total_tokens=150), ) - with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: - mock_aget.return_value = mock_response - await check_responses_cost_instance.check_responses_cost() + mock_llm_router.aget_responses = AsyncMock(return_value=mock_response) - assert mock_aget.await_count == 2 + await check_responses_cost_instance.check_responses_cost() + + assert mock_llm_router.aget_responses.await_count == 2 assert _completed_job_ids(mock_prisma_client) == [ "job-persist-fails", "job-persist-works", @@ -1194,16 +1201,11 @@ class TestCheckResponsesCost: is no handle on the generation. """ from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper - from litellm.responses.utils import ResponsesAPIRequestUtils from litellm.types.utils import SpecialEnums monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-key-for-response-ids") - provider_response_id = ResponsesAPIRequestUtils._build_responses_api_response_id( - custom_llm_provider="openai", - model_id="deployment-xyz", - response_id="resp_upstream_stable", - ) + provider_response_id = _routed_response_id("resp_upstream_stable") stale_advertised_id = "resp_" + str( encrypt_value_helper( value=SpecialEnums.LITELLM_MANAGED_RESPONSE_API_RESPONSE_ID_COMPLETE_STR.value.format( @@ -1243,16 +1245,11 @@ class TestCheckResponsesCost: ): """Rows created earlier carry the encrypted advertised id in both columns.""" from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper - from litellm.responses.utils import ResponsesAPIRequestUtils from litellm.types.utils import SpecialEnums monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-key-for-response-ids") - provider_response_id = ResponsesAPIRequestUtils._build_responses_api_response_id( - custom_llm_provider="openai", - model_id="deployment-xyz", - response_id="resp_legacy_upstream", - ) + provider_response_id = _routed_response_id("resp_legacy_upstream") legacy_id = "resp_" + str( encrypt_value_helper( value=SpecialEnums.LITELLM_MANAGED_RESPONSE_API_RESPONSE_ID_COMPLETE_STR.value.format( From 2aac6e109dbad841bcf4cbe6a01b06de70c89925 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Wed, 9 Sep 2026 18:35:09 -0700 Subject: [PATCH 07/13] refactor(responses): cut the cost poller's docstrings back to the non-obvious why The poller narrated its own straightforward behavior in seven multi-paragraph docstrings, which the repo's comment policy rules out. Each is now the claim a reader needs to avoid a wrong edit and nothing more. Also drops the deprecated `Dict` and `Optional` aliases the file still used. Claude-Session: https://claude.ai/code/session_01RHAjRxNhXTpKHeGMZ1nDKi --- .../common_utils/check_responses_cost.py | 78 ++++++------------- .../test_check_responses_cost.py | 6 +- 2 files changed, 25 insertions(+), 59 deletions(-) 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 47a6f38e8cb..30a88c04013 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py @@ -6,7 +6,7 @@ same route are non-inference and free. """ from datetime import datetime, timedelta, timezone -from typing import TYPE_CHECKING, Dict, Final, Optional, cast +from typing import TYPE_CHECKING, Final, cast import litellm from litellm._logging import verbose_proxy_logger @@ -48,18 +48,15 @@ class CheckResponsesCost: async def _get_response( self, response_id: str, - litellm_metadata: Dict[str, str], + litellm_metadata: dict[str, str], ) -> ResponsesAPIResponse: - """Fetch the upstream response, using deployment credentials when available. + """Fetch the upstream response through the deployment that served it. - LiteLLM-encoded response IDs carry the ``model_id`` of the deployment that - served the original request, so routing through ``llm_router`` applies that - deployment's ``api_base`` / ``api_key`` / ``api_version``, exactly like - ``GET /v1/responses/{id}`` does. ``litellm.aget_responses`` on its own only - sees provider env vars, so it fails for every deployment whose credentials - live in the config; the row then never leaves ``queued``. + A LiteLLM-encoded id carries its deployment's ``model_id``, so the router applies that + deployment's credentials. ``litellm.aget_responses`` only sees provider env vars, so a + config-only deployment's rows never leave ``queued``. """ - model_id: Optional[str] = ResponsesAPIRequestUtils.get_model_id_from_response_id(response_id) + model_id: str | None = ResponsesAPIRequestUtils.get_model_id_from_response_id(response_id) if model_id is None or self.llm_router.get_deployment(model_id=model_id) is None: return await litellm.aget_responses(response_id=response_id, litellm_metadata=litellm_metadata) router_response = await self.llm_router.aget_responses( @@ -70,15 +67,9 @@ class CheckResponsesCost: async def _expire_stale_rows( self, cutoff: datetime, batch_size: int ) -> int: - """Execute the bounded UPDATE that marks stale rows as 'stale_expired'. + """Run the bounded UPDATE that marks stale rows 'stale_expired'. - Isolated so it can be swapped / mocked in tests without touching the - orchestration logic in ``_cleanup_stale_managed_objects``. - - Uses PostgreSQL syntax (``$1::timestamptz``, ``LIMIT``, double-quoted - identifiers) which is the only dialect the proxy supports — every - ``schema.prisma`` in the repo sets ``provider = "postgresql"``. - Same pattern as ``spend_log_cleanup.py``. + PostgreSQL is the only dialect the proxy supports. Same pattern as ``spend_log_cleanup.py``. """ return await self.prisma_client.db.execute_raw( """ @@ -98,15 +89,9 @@ class CheckResponsesCost: ) async def _cleanup_stale_managed_objects(self) -> None: - """ - Mark managed objects older than MANAGED_OBJECT_STALENESS_CUTOFF_DAYS days - in non-terminal states as 'stale_expired'. These will never complete and - should not be polled. + """Retire rows stuck in a non-terminal state past the staleness cutoff, so they stop being polled. - Runs as a single DB query with a subquery LIMIT so no rows are loaded - into Python memory. Processes at most STALE_OBJECT_CLEANUP_BATCH_SIZE - rows per invocation to avoid overwhelming the DB when there is a large - backlog. + One query with a subquery LIMIT, so a large backlog never lands in Python memory. """ cutoff = datetime.now(timezone.utc) - timedelta(days=MANAGED_OBJECT_STALENESS_CUTOFF_DAYS) result = await self._expire_stale_rows(cutoff, STALE_OBJECT_CLEANUP_BATCH_SIZE) @@ -122,20 +107,12 @@ class CheckResponsesCost: return "batch_processed" in message or "unknown column" in message or "does not exist" in message async def _claim_job_for_costing(self, job: "LiteLLM_ManagedObjectTable") -> bool: - """Atomically flip batch_processed from false to true, returning whether this pod won the row. + """Atomically flip batch_processed false to true, returning whether this pod won the row. - Every pod and uvicorn worker schedules its own CheckResponsesCost against the shared table, - so without this compare-and-swap two of them select the same queued response in one window - and both bill it. The claim is taken before the read because the read is what prices the - job: ``aget_responses`` stamped with the poll origin writes the spend log itself, so there - is no later point at which to serialize. Schemas without the column can't be claimed, so - they keep the pre-existing behavior rather than silently billing nothing. - - A pod that dies between winning the claim and billing would otherwise strand the row: - it holds a claim nobody will release, and its status never reaches terminal, so every - later cycle re-selects it and loses. The ``updated_at`` arm takes such a claim back once - it has gone unbilled for longer than any live cycle could hold it. ``updated_at`` is - ``@updatedAt``, so a healthy in-flight claim refreshed moments ago is never stolen. + 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. """ abandoned_before: Final = datetime.now(timezone.utc) - timedelta( seconds=CLAIM_ABANDONED_AFTER_POLL_CYCLES * PROXY_BATCH_POLLING_INTERVAL @@ -162,10 +139,9 @@ class CheckResponsesCost: return claimed > 0 async def _release_job_claim(self, job: "LiteLLM_ManagedObjectTable") -> None: - """Give a claimed row back when the read did not bill it, so a later poll cycle retries it. + """Give a claimed row back when the read did not bill it, so a later cycle retries it. - A response still queued at the provider, or whose read raised, has no spend to record yet. - Holding the claim would retire it permanently, which is the failure #37050 hit on batches. + Holding the claim would retire the row unbilled, which is the failure #37050 hit on batches. """ try: await self.prisma_client.db.litellm_managedobjecttable.update_many( @@ -181,12 +157,8 @@ class CheckResponsesCost: 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. - Only ``status`` is written. The generation's usage and spend already land in - ``LiteLLM_SpendLogs`` unconditionally, so copying the response body onto this row would - duplicate content the provider still serves, on a table nothing ever deletes from. - - ``status`` stays the literal "completed" for every terminal provider status, matching - what this poller has always written, so stale-row expiry keeps skipping these rows. + 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. """ try: await self.prisma_client.db.litellm_managedobjecttable.update_many( @@ -199,13 +171,9 @@ class CheckResponsesCost: ) async def check_responses_cost(self): - """ - Check if background responses are complete and track their cost. - - Get all status="queued" or "in_progress" and file_purpose="response" jobs - - Query the provider to check if response is complete - - Cost is tracked by the get-responses call, billed because the poll is stamped - with BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN - - Mark responses in a terminal state as complete in the database + """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. """ try: await self._cleanup_stale_managed_objects() diff --git a/tests/proxy_unit_tests/test_check_responses_cost.py b/tests/proxy_unit_tests/test_check_responses_cost.py index 1e7ca7dd85d..4851669ff2a 100644 --- a/tests/proxy_unit_tests/test_check_responses_cost.py +++ b/tests/proxy_unit_tests/test_check_responses_cost.py @@ -843,10 +843,8 @@ class TestCheckResponsesCost: async def test_claim_is_taken_back_from_a_pod_that_died_holding_it( self, check_responses_cost_instance, mock_prisma_client, mock_llm_router ): - """A pod that dies between claiming and billing releases nothing, and the row's status - never reaches terminal, so without a lease every later cycle re-selects it and loses. - The window has to be longer than a live cycle can hold a claim and short enough that the - row is retried well before stale expiry gives up on it unbilled.""" + """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 cacfc47089cb0fd034ea875e018dbfd48a541fc8 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Wed, 9 Sep 2026 18:46:58 -0700 Subject: [PATCH 08/13] refactor(responses): give ResponsesIDSecurity a public provider_response_id Two callers only want the provider's own id behind an advertised one, and both reached past the class to get it out of the decrypt tuple. The poller also typed its rows as the pydantic projection, which has no `id`, while it is handed a Prisma row and reads `job.id` six times. Claude-Session: https://claude.ai/code/session_01RHAjRxNhXTpKHeGMZ1nDKi --- .../common_utils/check_responses_cost.py | 5 +- litellm/proxy/hooks/responses_id_security.py | 5 ++ .../proxy/response_api_endpoints/endpoints.py | 2 +- .../test_responses_id_security.py | 47 +++++++++++++++++++ 4 files changed, 56 insertions(+), 3 deletions(-) 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 30a88c04013..ec00e5cd03a 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py @@ -22,7 +22,8 @@ from litellm.types.llms.openai import ResponsesAPIResponse from litellm.types.utils import BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN if TYPE_CHECKING: - from litellm.proxy._types import LiteLLM_ManagedObjectTable + from prisma.models import LiteLLM_ManagedObjectTable + from litellm.proxy.utils import PrismaClient, ProxyLogging from litellm.router import Router @@ -207,7 +208,7 @@ class CheckResponsesCost: model_name = stored_response.get("model", None) # Decrypts rows written before model_object_id held the provider's own id. - responses_id_security, _, _ = ResponsesIDSecurity()._decrypt_response_id(job.model_object_id) + responses_id_security = ResponsesIDSecurity().provider_response_id(job.model_object_id) # Prepare metadata with model information for cost tracking litellm_metadata = { diff --git a/litellm/proxy/hooks/responses_id_security.py b/litellm/proxy/hooks/responses_id_security.py index 7e7f70d6f7e..b499006b64d 100644 --- a/litellm/proxy/hooks/responses_id_security.py +++ b/litellm/proxy/hooks/responses_id_security.py @@ -227,6 +227,11 @@ class ResponsesIDSecurity(CustomLogger): return True return False + def provider_response_id(self, response_id: str) -> str: + """The provider's own id behind an advertised one, returned unchanged when it is not encrypted.""" + original_response_id: Final = self._decrypt_response_id(response_id)[0] + return original_response_id + def _decrypt_response_id(self, response_id: str) -> tuple[str, str | None, str | None]: """ Returns: diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index 9f0e4937bba..38399f4bd21 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -67,7 +67,7 @@ async def store_background_response_object( ) return - provider_response_id, _, _ = ResponsesIDSecurity()._decrypt_response_id(response.id) + provider_response_id: Final = ResponsesIDSecurity().provider_response_id(response.id) await managed_files_obj.store_unified_object_id( unified_object_id=response.id, file_object=response, diff --git a/tests/test_litellm/test_responses_id_security.py b/tests/test_litellm/test_responses_id_security.py index a6081670172..44504194686 100644 --- a/tests/test_litellm/test_responses_id_security.py +++ b/tests/test_litellm/test_responses_id_security.py @@ -1094,3 +1094,50 @@ class TestClientSuppliedRetainedIdCannotBypassAuthorization: assert result["response_id"] == "resp_strangerownprovideridcccccccc" assert result["response_id"] != victim_provider_id + + +class TestProviderResponseId: + """The provider's own id behind an advertised one, for callers that only need that.""" + + def test_a_real_encrypted_id_round_trips_to_the_provider_id( + self, responses_id_security, monkeypatch + ): + from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-key-for-response-ids") + + advertised_id = "resp_" + str( + encrypt_value_helper( + value=SpecialEnums.LITELLM_MANAGED_RESPONSE_API_RESPONSE_ID_COMPLETE_STR.value.format( + "resp_provider_abc", "user-1", "team-1" + ) + ) + ) + + assert advertised_id != "resp_provider_abc" + assert responses_id_security.provider_response_id(advertised_id) == "resp_provider_abc" + + def test_two_encryptions_of_one_generation_resolve_to_the_same_provider_id( + self, responses_id_security, monkeypatch + ): + """Each advertised id carries a fresh nonce, so only the decrypted id can key a stored row.""" + from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-key-for-response-ids") + + payload = SpecialEnums.LITELLM_MANAGED_RESPONSE_API_RESPONSE_ID_COMPLETE_STR.value.format( + "resp_provider_abc", "user-1", "team-1" + ) + first = "resp_" + str(encrypt_value_helper(value=payload)) + second = "resp_" + str(encrypt_value_helper(value=payload)) + + assert first != second + assert responses_id_security.provider_response_id(first) == "resp_provider_abc" + assert responses_id_security.provider_response_id(second) == "resp_provider_abc" + + def test_a_raw_provider_id_is_returned_unchanged(self, responses_id_security, monkeypatch): + """Rows written before the provider id was stored hold an encrypted id, so both shapes + have to survive the same call.""" + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-key-for-response-ids") + + assert responses_id_security.provider_response_id("resp_provider_abc") == "resp_provider_abc" From 9a1d574493104cf716e468cec2924677aaea8e77 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Wed, 9 Sep 2026 18:57:55 -0700 Subject: [PATCH 09/13] refactor(responses): name the background row store by protocol, not the enterprise class The seam's type annotation pulled `litellm_enterprise` into this module's import graph, which check_unsafe_enterprise_import rejects outside a try-except. A Protocol carrying the one method the seam calls types it without the import and drops the local one the cast needed too. Claude-Session: https://claude.ai/code/session_01RHAjRxNhXTpKHeGMZ1nDKi --- .../proxy/response_api_endpoints/endpoints.py | 30 +++++++++++++------ 1 file changed, 21 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index 38399f4bd21..65f6d403c46 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -4,7 +4,7 @@ import time from collections.abc import AsyncIterator, Awaitable, Mapping from enum import Enum from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, NamedTuple, Protocol, cast, get_args +from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, Protocol, cast, get_args from uuid import uuid4 import fastapi @@ -39,16 +39,32 @@ from litellm.types.responses.main import DeleteResponseResult from litellm.types.utils import TokenCountResponse if TYPE_CHECKING: - from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles - from litellm.router import Router router: Final = APIRouter() +class BackgroundResponseStore(Protocol): + """The one managed-object write a queued background response needs. + + Naming it here keeps this module from importing the enterprise hook that implements it. + """ + + async def store_unified_object_id( + self, + unified_object_id: str, + file_object: ResponsesAPIResponse, + litellm_parent_otel_span: object | None, + model_object_id: str, + file_purpose: Literal["response"], + user_api_key_dict: UserAPIKeyAuth, + persist_attribution: bool = False, + ) -> None: ... + + async def store_background_response_object( response: ResponsesAPIResponse, - managed_files_obj: "_PROXY_LiteLLMManagedFiles", + managed_files_obj: BackgroundResponseStore, user_api_key_dict: UserAPIKeyAuth, ) -> None: """Record a queued background response so the cost poller can find and bill it. @@ -408,12 +424,8 @@ async def responses_api( and isinstance(response, ResponsesAPIResponse) and response.status in ("queued", "in_progress") ): - from litellm_enterprise.proxy.hooks.managed_files import ( - _PROXY_LiteLLMManagedFiles, - ) - managed_files_obj: Final = cast( - _PROXY_LiteLLMManagedFiles | None, + BackgroundResponseStore | None, proxy_logging_obj.get_proxy_hook("managed_files"), ) From 33d9464f25d8cf65e90de70f093c2fbf5ddf7a3e Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Wed, 9 Sep 2026 19:10:39 -0700 Subject: [PATCH 10/13] test(responses): cover the gate that decides a background row gets written The storage branch's three conditions sat inline in `responses_api`, so nothing proved a foreground create or an already-terminal one stays out of the managed table. They move into `should_store_background_response`, which the endpoint calls and the tests exercise across both arms. Claude-Session: https://claude.ai/code/session_01RHAjRxNhXTpKHeGMZ1nDKi --- .../proxy/response_api_endpoints/endpoints.py | 21 +++++-- .../response_api_endpoints/test_endpoints.py | 57 +++++++++++++++++++ 2 files changed, 72 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index 65f6d403c46..415d630e015 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -62,6 +62,20 @@ class BackgroundResponseStore(Protocol): ) -> None: ... +_STORABLE_BACKGROUND_STATUSES: Final[frozenset[str]] = frozenset({"queued", "in_progress"}) + + +def should_store_background_response(data: Mapping[str, object], response: object) -> bool: + """Whether a create just produced a generation the cost poller will have to bill later. + + Only a background create leaves usage unreported, and only while the provider has not + finished it; anything already terminal reported its usage on this very call. + """ + if not data.get("background") or not isinstance(response, ResponsesAPIResponse): + return False + return response.status in _STORABLE_BACKGROUND_STATUSES + + async def store_background_response_object( response: ResponsesAPIResponse, managed_files_obj: BackgroundResponseStore, @@ -418,12 +432,7 @@ async def responses_api( version=version, ) - # Store in managed objects table if background mode is enabled - if ( - data.get("background") - and isinstance(response, ResponsesAPIResponse) - and response.status in ("queued", "in_progress") - ): + if should_store_background_response(data, response): managed_files_obj: Final = cast( BackgroundResponseStore | None, proxy_logging_obj.get_proxy_hook("managed_files"), diff --git a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py index f528dc6ba58..907feb971dd 100644 --- a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py @@ -2069,3 +2069,60 @@ class TestBackgroundResponseManagedObjectId: store = await self._stored_kwargs(self._encrypted_id("resp_no_deployment"), model_id=None) store.assert_not_awaited() + + +class TestShouldStoreBackgroundResponse: + """The gate `responses_api` applies before it writes a managed row. + + Storing a foreground create would bill a generation whose usage the create already + reported, and storing one the provider has already finished leaves a row no poll can + retire, so both arms have to stay closed. + """ + + @staticmethod + def _response(status: str): + from litellm.types.llms.openai import ResponsesAPIResponse + + return ResponsesAPIResponse( + id="resp_abc", + created_at=0, + model="gpt-4o", + object="response", + output=[], + parallel_tool_calls=False, + tool_choice="auto", + tools=[], + status=status, + ) + + @pytest.mark.parametrize("status", ["queued", "in_progress"]) + def test_a_background_create_the_provider_has_not_finished_is_stored(self, status): + from litellm.proxy.response_api_endpoints.endpoints import ( + should_store_background_response, + ) + + assert should_store_background_response({"background": True}, self._response(status)) is True + + @pytest.mark.parametrize("status", ["completed", "failed", "cancelled", "incomplete"]) + def test_a_background_create_already_terminal_is_not_stored(self, status): + from litellm.proxy.response_api_endpoints.endpoints import ( + should_store_background_response, + ) + + assert should_store_background_response({"background": True}, self._response(status)) is False + + @pytest.mark.parametrize("data", [{}, {"background": False}, {"background": None}]) + def test_a_foreground_create_is_never_stored(self, data): + from litellm.proxy.response_api_endpoints.endpoints import ( + should_store_background_response, + ) + + assert should_store_background_response(data, self._response("queued")) is False + + def test_a_streaming_or_error_result_is_not_mistaken_for_a_response(self): + """The create path can hand back a streaming iterator, which has no status to read.""" + from litellm.proxy.response_api_endpoints.endpoints import ( + should_store_background_response, + ) + + assert should_store_background_response({"background": True}, object()) is False From 08ff09f5859bb025fd6baf7533a30a6ae4f00008 Mon Sep 17 00:00:00 2001 From: jesus Date: Tue, 15 Sep 2026 17:23:38 +0000 Subject: [PATCH 11/13] fix(responses): finalize the managed row before the billing read so a stolen claim cannot rebill Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../common_utils/check_responses_cost.py | 78 +++--- .../test_check_responses_cost.py | 240 ++++++++++++++++-- 2 files changed, 262 insertions(+), 56 deletions(-) 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", From cea5db16c41de34f081928ad90212e059388d740 Mon Sep 17 00:00:00 2001 From: jesus Date: Tue, 15 Sep 2026 17:37:31 +0000 Subject: [PATCH 12/13] fix(responses): avoid loop Final annotations in cost polling Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/common_utils/check_responses_cost.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) 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 cd52fa6e9ec..349fa873792 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py @@ -214,7 +214,7 @@ 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) - probe_metadata: Final = { + probe_metadata: dict[str, str] = { "user_api_key_user_id": job.created_by or "default-user-id", **({"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 {}), @@ -263,7 +263,7 @@ class CheckResponsesCost: verbose_proxy_logger.info( f"Response {unified_object_id} has terminal status {response.status}, marked as complete" ) - billing_metadata: Final = { + billing_metadata: dict[str, str] = { **probe_metadata, INTERNAL_CALL_ORIGIN_METADATA_KEY: BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN, } From 350ed9541173c7e65936bc975f54c1203980a531 Mon Sep 17 00:00:00 2001 From: jesus Date: Thu, 17 Sep 2026 22:30:47 +0000 Subject: [PATCH 13/13] fix(responses): bill background responses against the creating org and tags Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../common_utils/check_responses_cost.py | 12 +++- .../proxy/response_api_endpoints/endpoints.py | 10 ++- .../test_check_responses_cost.py | 67 +++++++++++++++++++ .../response_api_endpoints/test_endpoints.py | 27 +++++++- 4 files changed, 111 insertions(+), 5 deletions(-) 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 349fa873792..3f6cce87876 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py @@ -48,7 +48,7 @@ class CheckResponsesCost: async def _get_response( self, response_id: str, - litellm_metadata: dict[str, str], + litellm_metadata: dict[str, object], ) -> ResponsesAPIResponse: """Fetch the upstream response through the deployment that served it. @@ -214,11 +214,17 @@ 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) - probe_metadata: dict[str, str] = { + probe_metadata: dict[str, object] = { "user_api_key_user_id": job.created_by or "default-user-id", **({"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 {}), + **({"user_api_key_org_id": job.org_id} if job.org_id else {}), + **( + {"tags": [tag for tag in job.request_tags if isinstance(tag, str)]} + if isinstance(job.request_tags, list) and job.request_tags + else {} + ), } except Exception as e: @@ -263,7 +269,7 @@ class CheckResponsesCost: verbose_proxy_logger.info( f"Response {unified_object_id} has terminal status {response.status}, marked as complete" ) - billing_metadata: dict[str, str] = { + billing_metadata: dict[str, object] = { **probe_metadata, INTERNAL_CALL_ORIGIN_METADATA_KEY: BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN, } diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index 415d630e015..02024dc14f4 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -1,7 +1,7 @@ import asyncio import json import time -from collections.abc import AsyncIterator, Awaitable, Mapping +from collections.abc import AsyncIterator, Awaitable, Mapping, Sequence from enum import Enum from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, Protocol, cast, get_args @@ -30,6 +30,9 @@ from litellm.proxy.common_utils.http_parsing_utils import ( _read_request_body, _safe_set_request_parsed_body, ) +from litellm.proxy.pass_through_endpoints.llm_provider_handlers.batch_attribution import ( + request_tags_from_metadata, +) from litellm.types.llms.openai import ( REASONING_EFFORT, ResponsesAPIOptionalRequestParams, @@ -58,6 +61,7 @@ class BackgroundResponseStore(Protocol): model_object_id: str, file_purpose: Literal["response"], user_api_key_dict: UserAPIKeyAuth, + request_tags: Sequence[str] | None = None, persist_attribution: bool = False, ) -> None: ... @@ -80,6 +84,7 @@ async def store_background_response_object( response: ResponsesAPIResponse, managed_files_obj: BackgroundResponseStore, user_api_key_dict: UserAPIKeyAuth, + data: Mapping[str, object], ) -> None: """Record a queued background response so the cost poller can find and bill it. @@ -98,6 +103,7 @@ async def store_background_response_object( return provider_response_id: Final = ResponsesIDSecurity().provider_response_id(response.id) + litellm_metadata: Final = data.get("litellm_metadata") await managed_files_obj.store_unified_object_id( unified_object_id=response.id, file_object=response, @@ -105,6 +111,7 @@ async def store_background_response_object( model_object_id=provider_response_id, file_purpose="response", user_api_key_dict=user_api_key_dict, + request_tags=request_tags_from_metadata(litellm_metadata if isinstance(litellm_metadata, dict) else {}), persist_attribution=True, ) verbose_proxy_logger.info("Stored background response %s in managed objects table", response.id) @@ -444,6 +451,7 @@ async def responses_api( response=response, managed_files_obj=managed_files_obj, user_api_key_dict=user_api_key_dict, + data=data, ) except Exception as e: verbose_proxy_logger.error("Failed to store background response in managed objects table: %s", e) diff --git a/tests/proxy_unit_tests/test_check_responses_cost.py b/tests/proxy_unit_tests/test_check_responses_cost.py index e6475c8b452..e6136f5491d 100644 --- a/tests/proxy_unit_tests/test_check_responses_cost.py +++ b/tests/proxy_unit_tests/test_check_responses_cost.py @@ -1049,6 +1049,73 @@ class TestCheckResponsesCost: } assert _release_calls(mock_prisma_client) == [] + @pytest.mark.asyncio + async def test_billing_read_includes_managed_row_attribution( + 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 + + attributed_job = MagicMock() + attributed_job.unified_object_id = "resp_attributed" + attributed_job.model_object_id = _routed_response_id("resp_attributed") + attributed_job.created_by = "test-user" + attributed_job.org_id = "org-1" + attributed_job.request_tags = ["tag-a", "tag-b"] + attributed_job.id = "job-attributed" + attributed_job.file_object = {"model": "gpt-5", "id": "resp_attributed"} + + unattributed_job = MagicMock() + unattributed_job.unified_object_id = "resp_unattributed" + unattributed_job.model_object_id = _routed_response_id("resp_unattributed") + unattributed_job.created_by = "test-user" + unattributed_job.org_id = None + unattributed_job.request_tags = [] + unattributed_job.id = "job-unattributed" + unattributed_job.file_object = {"model": "gpt-5", "id": "resp_unattributed"} + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[attributed_job, unattributed_job] + ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) + mock_llm_router.aget_responses = AsyncMock( + side_effect=[ + ResponsesAPIResponse( + id="resp_attributed", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + ] + * 2 + + [ + ResponsesAPIResponse( + id="resp_unattributed", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + ] + * 2 + ) + + await check_responses_cost_instance.check_responses_cost() + + billing_metadata = [ + call.kwargs["litellm_metadata"] + for call in mock_llm_router.aget_responses.await_args_list + if call.kwargs["litellm_metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) + == BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN + ] + assert billing_metadata[0]["user_api_key_org_id"] == "org-1" + assert billing_metadata[0]["tags"] == ["tag-a", "tag-b"] + assert "user_api_key_org_id" not in billing_metadata[1] + assert "tags" not in billing_metadata[1] + @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 diff --git a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py index 907feb971dd..aa5287cf7ab 100644 --- a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py @@ -1997,7 +1997,12 @@ class TestBackgroundResponseManagedObjectId: response._hidden_params = {"model_id": model_id} if model_id else {} return response - async def _stored_kwargs(self, advertised_id: str, model_id: str | None = "deployment-1"): + async def _stored_kwargs( + self, + advertised_id: str, + model_id: str | None = "deployment-1", + data: dict[str, object] | None = None, + ): from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.response_api_endpoints.endpoints import ( store_background_response_object, @@ -2010,6 +2015,7 @@ class TestBackgroundResponseManagedObjectId: response=self._queued_response(advertised_id, model_id), managed_files_obj=managed_files_obj, user_api_key_dict=UserAPIKeyAuth(api_key="sk-1234", user_id="u-1", team_id="t-1"), + data=data if data is not None else {"background": True}, ) return managed_files_obj.store_unified_object_id @@ -2070,6 +2076,25 @@ class TestBackgroundResponseManagedObjectId: store.assert_not_awaited() + @pytest.mark.asyncio + async def test_request_tags_are_forwarded_from_litellm_metadata(self, monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-regression-salt") + + tagged_store = await self._stored_kwargs( + self._encrypted_id("resp_tagged"), + data={ + "background": True, + "litellm_metadata": {"tags": ["tag-a", "tag-b"]}, + }, + ) + untagged_store = await self._stored_kwargs( + self._encrypted_id("resp_untagged"), + data={"background": True}, + ) + + assert tagged_store.await_args.kwargs["request_tags"] == ("tag-a", "tag-b") + assert untagged_store.await_args.kwargs["request_tags"] is None + class TestShouldStoreBackgroundResponse: """The gate `responses_api` applies before it writes a managed row.