This commit is contained in:
devin-ai-integration[bot] 2026-09-23 15:00:12 -04:00 • committed by GitHub
commit d5902fe79a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
20 changed files with 1403 additions and 301 deletions

View file

@ -1,12 +1,11 @@
"""
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
from typing import TYPE_CHECKING, Dict, Final, Optional, Protocol, cast
from typing import TYPE_CHECKING, Final, cast
import litellm
from litellm._logging import verbose_proxy_logger
@ -14,6 +13,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,30 +21,14 @@ from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.utils import BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN
if TYPE_CHECKING:
from prisma.models import LiteLLM_ManagedObjectTable
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.repositories.prisma_protocols import TableActions
from litellm.router import Router
TERMINAL_RESPONSE_STATUSES = frozenset({"completed", "failed", "cancelled", "incomplete"})
class _ManagedObjectRow(Protocol):
@property
def id(self) -> str: ...
@property
def unified_object_id(self) -> str: ...
@property
def created_by(self) -> str | None: ...
@property
def file_object(self) -> object: ...
def _managed_object_table(prisma_client: "PrismaClient") -> "TableActions[_ManagedObjectRow]":
table: Final[TableActions[_ManagedObjectRow]] = prisma_client.db.litellm_managedobjecttable
return table
CLAIM_ABANDONED_AFTER_POLL_CYCLES: Final = 3
class CheckResponsesCost:
@ -64,18 +48,15 @@ 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, 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(
@ -86,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(
"""
@ -114,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)
@ -132,14 +101,84 @@ class CheckResponsesCost:
f"(older than {MANAGED_OBJECT_STALENESS_CUTOFF_DAYS} days) as stale_expired"
)
async def check_responses_cost(self):
@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 false to true, returning whether this pod won the row.
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.
"""
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
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 cycle retries it.
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(
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 _mark_job_completed(self, job: "LiteLLM_ManagedObjectTable") -> bool:
"""CAS the durable billed marker before the billing read, returning whether this pod won.
Only queued or in-progress rows may transition to the literal "completed" value.
"""
try:
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 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()
@ -148,7 +187,7 @@ class CheckResponsesCost:
f"CheckResponsesCost: stale cleanup failed (poll will continue): {cleanup_err}"
)
jobs = await _managed_object_table(self.prisma_client).find_many(
jobs = await self.prisma_client.db.litellm_managedobjecttable.find_many(
where={
"status": {"in": ["queued", "in_progress"]},
"file_purpose": "response",
@ -158,7 +197,7 @@ class CheckResponsesCost:
)
verbose_proxy_logger.debug(f"Found {len(jobs)} response jobs to check")
completed_jobs: Final[list[_ManagedObjectRow]] = []
completed_count: int = 0
for job in jobs:
unified_object_id = job.unified_object_id
@ -172,48 +211,77 @@ 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().provider_response_id(job.model_object_id)
# Prepare metadata with model information for cost tracking
litellm_metadata = {
probe_metadata: dict[str, object] = {
"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 {}),
**({"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 {}
),
}
# 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
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 _managed_object_table(self.prisma_client).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=probe_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
if not await self._mark_job_completed(job):
verbose_proxy_logger.debug(
f"Response {unified_object_id} was already finalized by another poller"
)
continue
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: dict[str, object] = {
**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")

View file

@ -328,11 +328,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

@ -1688,7 +1688,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"}
@ -1766,9 +1766,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
@ -2543,9 +2541,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

@ -1836,7 +1836,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
@ -6444,7 +6444,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

@ -2890,7 +2890,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

@ -226,6 +226,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:

View file

@ -6,7 +6,7 @@ from collections.abc import AsyncIterator, Awaitable, Mapping, Sequence
from enum import Enum
from functools import partial
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
@ -34,6 +34,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.proxy.route_llm_request import raise_if_required_body_param_missing
from litellm.types.llms.openai import (
REASONING_EFFORT,
@ -48,6 +51,77 @@ if TYPE_CHECKING:
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,
request_tags: Sequence[str] | None = None,
persist_attribution: bool = False,
) -> 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,
user_api_key_dict: UserAPIKeyAuth,
data: Mapping[str, object],
) -> 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: 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,
litellm_parent_otel_span=None,
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)
_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
@ -372,49 +446,22 @@ async def responses_api(
version=version,
)
# 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 should_store_background_response(data, response):
managed_files_obj: Final = cast(
BackgroundResponseStore | 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:
# 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")
# 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,
file_purpose="response",
user_api_key_dict=user_api_key_dict,
)
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,
data=data,
)
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:

View file

@ -507,7 +507,7 @@ def get_logging_payload(
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):

File diff suppressed because it is too large Load diff

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

@ -6767,14 +6767,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)
@ -6782,7 +6782,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(
@ -367,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(
@ -396,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
@ -427,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(
@ -451,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
@ -480,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(
@ -504,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
@ -533,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,
@ -551,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

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

@ -6397,15 +6397,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 (
@ -6428,7 +6427,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

@ -2427,3 +2427,195 @@ 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 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 handle on the generation 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)}"
@staticmethod
def _queued_response(advertised_id: str, model_id: str | None = "deployment-1"):
from litellm.types.llms.openai import ResponsesAPIResponse
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": model_id} if model_id else {}
return response
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,
)
managed_files_obj = MagicMock()
managed_files_obj.store_unified_object_id = AsyncMock()
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"),
data=data if data is not None else {"background": True},
)
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"
)
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"] == advertised_id
assert kwargs["file_object"].id == advertised_id
@pytest.mark.asyncio
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._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"
@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."""
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-regression-salt")
store = await self._stored_kwargs("resp_rawprovider999")
assert store.await_args.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()
@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.
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

View file

@ -3948,13 +3948,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

@ -6426,7 +6426,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),
@ -6434,7 +6434,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):

View file

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