mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-28 01:32:17 +00:00
Merge 38cc3e6715 into 9fd25b2228
This commit is contained in:
commit
d5902fe79a
20 changed files with 1403 additions and 301 deletions
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 = (
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue