From bb81a9f9f173747907a9a354af058e410a7679b0 Mon Sep 17 00:00:00 2001 From: jesus Date: Tue, 8 Sep 2026 15:06:37 +0000 Subject: [PATCH] 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):