mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
b67137b67f
commit
bb81a9f9f1
17 changed files with 68 additions and 72 deletions
|
|
@ -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"
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 = (
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue