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>
This commit is contained in:
jesus 2026-09-08 15:06:37 +00:00
parent b67137b67f
commit bb81a9f9f1
17 changed files with 68 additions and 72 deletions

View file

@ -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"
)

View file

@ -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

View file

@ -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:

View file

@ -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)

View file

@ -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:

View file

@ -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 = (

View file

@ -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)

View file

@ -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(

View file

@ -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):

View file

@ -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

View file

@ -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():

View file

@ -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
)

View file

@ -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(

View file

@ -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."""

View file

@ -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

View file

@ -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():

View file

@ -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):