diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py index cdeea0d3d4b..3f6cce87876 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py @@ -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") diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index 5ac7c1e53c1..d73a6b7e5d1 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -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 diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index 749f0ce4fcb..27b295e55d9 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -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: diff --git a/litellm/integrations/otel/plumbing/metrics.py b/litellm/integrations/otel/plumbing/metrics.py index e1623f4697f..7bc080c36ee 100644 --- a/litellm/integrations/otel/plumbing/metrics.py +++ b/litellm/integrations/otel/plumbing/metrics.py @@ -224,7 +224,7 @@ class GenAIMetricRecorder: common_attrs: Final = self._filter_attributes(self._bounded_attributes(kwargs)) duration_s: Final = (end_time - start_time).total_seconds() usage_is_replayed: Final = is_unbilled_non_inference_call_from_params( - kwargs.get("call_type"), kwargs.get("litellm_params"), response_obj + kwargs.get("call_type"), kwargs.get("litellm_params") ) self._metrics.operation_duration.record(duration_s, attributes=common_attrs) diff --git a/litellm/litellm_core_utils/internal_call_metadata.py b/litellm/litellm_core_utils/internal_call_metadata.py index 87f007ca1d5..c4ab791a0d3 100644 --- a/litellm/litellm_core_utils/internal_call_metadata.py +++ b/litellm/litellm_core_utils/internal_call_metadata.py @@ -54,35 +54,18 @@ budget-checked like the request that spawned it. Everything else on the parent's be a lie on a sub-call that runs after it returned.""" -def is_background_response(response: object) -> bool: - """Whether a retrieved object is a response created with ``background=true``. - - Such a create returns ``status="queued"`` and no usage at all, so nothing has billed the - job by the time anyone reads it back. Accepts the response as a mapping or a model, - because the callers hold it in both shapes. - """ - if isinstance(response, Mapping): - return response.get("background") is True - return getattr(response, "background", None) is True - - def is_unbilled_non_inference_call( call_type: str | None, metadata: Mapping[str, object] | None, - response: object, ) -> bool: - """A read/management route priced at zero, because the usage it reports belongs to the - call that created the object it just read. + """Reads of stored objects are priced at zero because their usage belongs to the create. - Retrieving a background response is the exception, and the enterprise cost poller's read - is the same exception seen from the other side: that job's create billed nothing, so its - retrieval is the only place the spend is ever visible. Pricing those at zero would lose - the spend rather than deduplicate it. + A background create bills nothing, so the enterprise cost poller's read stamped with + ``BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN`` prices normally and marks the managed-object + row completed so it prices once. """ if call_type not in NON_INFERENCE_CALL_TYPES: return False - if is_background_response(response): - return False if metadata is None: return True return metadata.get(INTERNAL_CALL_ORIGIN_METADATA_KEY) != BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN @@ -91,7 +74,6 @@ def is_unbilled_non_inference_call( def is_unbilled_non_inference_call_from_params( call_type: str | None, litellm_params: Mapping[str, object] | None, - response: object, ) -> bool: """:func:`is_unbilled_non_inference_call` for callers holding raw ``litellm_params``. @@ -105,7 +87,7 @@ def is_unbilled_non_inference_call_from_params( metadata: Final = ( StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params) if litellm_params is not None else None ) - return is_unbilled_non_inference_call(call_type, metadata, response) + return is_unbilled_non_inference_call(call_type, metadata) def sanitize_user_api_key_auth(auth: object) -> object: diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 8e28a0d543d..55698c1ac84 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -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 = ( diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index c40090233be..e328ff262d8 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -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) diff --git a/litellm/proxy/hooks/responses_id_security.py b/litellm/proxy/hooks/responses_id_security.py index bdf7e2ab53d..e069d4a79a6 100644 --- a/litellm/proxy/hooks/responses_id_security.py +++ b/litellm/proxy/hooks/responses_id_security.py @@ -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: diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index 36b7a3a4a8a..6d48d19d151 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -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: diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index 8f85ecdd480..c43d942ae6f 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -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): diff --git a/tests/proxy_unit_tests/test_check_responses_cost.py b/tests/proxy_unit_tests/test_check_responses_cost.py index e806e9a3394..e6136f5491d 100644 --- a/tests/proxy_unit_tests/test_check_responses_cost.py +++ b/tests/proxy_unit_tests/test_check_responses_cost.py @@ -2,9 +2,8 @@ Unit tests for CheckResponsesCost class """ -import asyncio -from datetime import datetime -from unittest.mock import AsyncMock, MagicMock, Mock, patch +from datetime import datetime, timedelta, timezone +from unittest.mock import AsyncMock, MagicMock, call, patch import pytest @@ -12,6 +11,47 @@ from litellm.constants import MAX_OBJECTS_PER_POLL_CYCLE from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse +def _update_many_calls_writing(mock_prisma_client, matches_data): + return [ + call + for call in mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list + if matches_data(call.kwargs["data"]) + ] + + +def _completion_calls(mock_prisma_client): + return _update_many_calls_writing( + mock_prisma_client, lambda data: data.get("status") == "completed" + ) + + +def _completed_job_ids(mock_prisma_client): + return [call.kwargs["where"]["id"] for call in _completion_calls(mock_prisma_client)] + + +def _claim_calls(mock_prisma_client): + return _update_many_calls_writing( + mock_prisma_client, lambda data: data == {"batch_processed": True} + ) + + +def _release_calls(mock_prisma_client): + return _update_many_calls_writing( + mock_prisma_client, lambda data: data == {"batch_processed": False} + ) + + +def _routed_response_id(provider_response_id): + """A LiteLLM-encoded id names a deployment, which is what sends the poll's read through the router.""" + from litellm.responses.utils import ResponsesAPIRequestUtils + + return ResponsesAPIRequestUtils._build_responses_api_response_id( + custom_llm_provider="openai", + model_id="deployment-xyz", + response_id=provider_response_id, + ) + + class TestCheckResponsesCost: """Test suite for CheckResponsesCost class""" @@ -112,6 +152,7 @@ class TestCheckResponsesCost: # Mock job with response ID mock_job = MagicMock() mock_job.unified_object_id = "resp_test_123" + mock_job.model_object_id = "resp_test_123" mock_job.created_by = "test-user" mock_job.id = "job-123" mock_job.file_object = {"model": "gpt-4o", "id": "resp_test_123"} @@ -135,7 +176,7 @@ class TestCheckResponsesCost: ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) # Run the check with mocked litellm.aget_responses @@ -144,14 +185,8 @@ class TestCheckResponsesCost: await check_responses_cost_instance.check_responses_cost() - # update_many should only contain the job completion call - calls = ( - mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list - ) - assert len(calls) == 1 - completion_call = calls[0] - assert completion_call[1]["data"]["status"] == "completed" - assert completion_call[1]["where"]["id"]["in"] == ["job-123"] + assert _completed_job_ids(mock_prisma_client) == ["job-123"] + assert _release_calls(mock_prisma_client) == [] @pytest.mark.asyncio async def test_check_responses_cost_with_failed_response( @@ -161,6 +196,7 @@ class TestCheckResponsesCost: # Mock job mock_job = MagicMock() mock_job.unified_object_id = "resp_test_456" + mock_job.model_object_id = "resp_test_456" mock_job.created_by = "test-user" mock_job.id = "job-456" mock_job.file_object = {"model": "gpt-4o", "id": "resp_test_456"} @@ -180,7 +216,7 @@ class TestCheckResponsesCost: ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) # Run the check @@ -189,12 +225,8 @@ class TestCheckResponsesCost: await check_responses_cost_instance.check_responses_cost() - # update_many should only contain the job completion call - calls = ( - mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list - ) - assert len(calls) == 1 - assert calls[0][1]["data"]["status"] == "completed" + assert _completed_job_ids(mock_prisma_client) == ["job-456"] + assert _release_calls(mock_prisma_client) == [] @pytest.mark.asyncio async def test_check_responses_cost_with_cancelled_response( @@ -204,6 +236,7 @@ class TestCheckResponsesCost: # Mock job mock_job = MagicMock() mock_job.unified_object_id = "resp_test_789" + mock_job.model_object_id = "resp_test_789" mock_job.created_by = "test-user" mock_job.id = "job-789" mock_job.file_object = {"model": "gpt-4o", "id": "resp_test_789"} @@ -223,7 +256,7 @@ class TestCheckResponsesCost: ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) # Run the check @@ -232,12 +265,8 @@ class TestCheckResponsesCost: await check_responses_cost_instance.check_responses_cost() - # update_many should only contain the job completion call - calls = ( - mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list - ) - assert len(calls) == 1 - assert calls[0][1]["data"]["status"] == "completed" + assert _completed_job_ids(mock_prisma_client) == ["job-789"] + assert _release_calls(mock_prisma_client) == [] @pytest.mark.asyncio async def test_check_responses_cost_with_in_progress_response( @@ -247,6 +276,7 @@ class TestCheckResponsesCost: # Mock job mock_job = MagicMock() mock_job.unified_object_id = "resp_test_in_progress" + mock_job.model_object_id = "resp_test_in_progress" mock_job.created_by = "test-user" mock_job.id = "job-in-progress" mock_job.file_object = {"model": "gpt-4o", "id": "resp_test_in_progress"} @@ -266,7 +296,7 @@ class TestCheckResponsesCost: ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) # Run the check @@ -276,10 +306,7 @@ class TestCheckResponsesCost: await check_responses_cost_instance.check_responses_cost() # No job completion update_many — response is still in progress - calls = ( - mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list - ) - assert len(calls) == 0 + assert _completion_calls(mock_prisma_client) == [] # Stale cleanup still ran via _expire_stale_rows check_responses_cost_instance._expire_stale_rows.assert_called_once() @@ -291,6 +318,7 @@ class TestCheckResponsesCost: # Mock job mock_job = MagicMock() mock_job.unified_object_id = "resp_test_queued" + mock_job.model_object_id = "resp_test_queued" mock_job.created_by = "test-user" mock_job.id = "job-queued" mock_job.file_object = {"model": "gpt-4o", "id": "resp_test_queued"} @@ -310,7 +338,7 @@ class TestCheckResponsesCost: ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) # Run the check @@ -320,10 +348,7 @@ class TestCheckResponsesCost: await check_responses_cost_instance.check_responses_cost() # No job completion update_many — response is still queued - calls = ( - mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list - ) - assert len(calls) == 0 + assert _completion_calls(mock_prisma_client) == [] # Stale cleanup still ran via _expire_stale_rows check_responses_cost_instance._expire_stale_rows.assert_called_once() @@ -335,6 +360,7 @@ class TestCheckResponsesCost: # Mock job mock_job = MagicMock() mock_job.unified_object_id = "resp_test_error" + mock_job.model_object_id = "resp_test_error" mock_job.created_by = "test-user" mock_job.id = "job-error" mock_job.file_object = {"model": "gpt-4o", "id": "resp_test_error"} @@ -344,7 +370,7 @@ class TestCheckResponsesCost: ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) # Run the check with mocked exception @@ -357,10 +383,7 @@ class TestCheckResponsesCost: await check_responses_cost_instance.check_responses_cost() # No job completion update_many — exception skipped the job - calls = ( - mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list - ) - assert len(calls) == 0 + assert _completion_calls(mock_prisma_client) == [] # Stale cleanup still ran via _expire_stale_rows check_responses_cost_instance._expire_stale_rows.assert_called_once() @@ -372,18 +395,21 @@ class TestCheckResponsesCost: # Mock multiple jobs mock_job1 = MagicMock() mock_job1.unified_object_id = "resp_test_1" + mock_job1.model_object_id = "resp_test_1" mock_job1.created_by = "user1" mock_job1.id = "job-1" mock_job1.file_object = {"model": "gpt-4o", "id": "resp_test_1"} mock_job2 = MagicMock() mock_job2.unified_object_id = "resp_test_2" + mock_job2.model_object_id = "resp_test_2" mock_job2.created_by = "user2" mock_job2.id = "job-2" mock_job2.file_object = {"model": "gpt-4o", "id": "resp_test_2"} mock_job3 = MagicMock() mock_job3.unified_object_id = "resp_test_3" + mock_job3.model_object_id = "resp_test_3" mock_job3.created_by = "user3" mock_job3.id = "job-3" mock_job3.file_object = {"model": "gpt-4o", "id": "resp_test_3"} @@ -429,25 +455,22 @@ class TestCheckResponsesCost: ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) # Run the check with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: - mock_aget.side_effect = [mock_response1, mock_response2, mock_response3] + mock_aget.side_effect = [ + mock_response1, + mock_response1, + mock_response2, + mock_response3, + mock_response3, + ] await check_responses_cost_instance.check_responses_cost() - # update_many should only contain the job completion call - calls = ( - mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list - ) - assert len(calls) == 1 - completion_call = calls[0] - assert len(completion_call[1]["where"]["id"]["in"]) == 2 - assert "job-1" in completion_call[1]["where"]["id"]["in"] - assert "job-3" in completion_call[1]["where"]["id"]["in"] - assert "job-2" not in completion_call[1]["where"]["id"]["in"] + assert _completed_job_ids(mock_prisma_client) == ["job-1", "job-3"] @pytest.mark.asyncio async def test_encoded_response_id_is_fetched_through_router( @@ -472,6 +495,7 @@ class TestCheckResponsesCost: mock_job = MagicMock() mock_job.unified_object_id = encoded_response_id + mock_job.model_object_id = encoded_response_id mock_job.created_by = "test-user" mock_job.id = "job-router" mock_job.file_object = {"model": "azure-gpt-5", "id": encoded_response_id} @@ -480,7 +504,7 @@ class TestCheckResponsesCost: return_value=[mock_job] ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) mock_llm_router.aget_responses = AsyncMock( @@ -511,12 +535,7 @@ class TestCheckResponsesCost: == encoded_response_id ) - calls = ( - mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list - ) - assert len(calls) == 1 - assert calls[0][1]["data"]["status"] == "completed" - assert calls[0][1]["where"]["id"]["in"] == ["job-router"] + assert _completed_job_ids(mock_prisma_client) == ["job-router"] @pytest.mark.asyncio async def test_encrypted_response_id_is_fetched_through_router( @@ -548,6 +567,7 @@ class TestCheckResponsesCost: mock_job = MagicMock() mock_job.unified_object_id = encrypted_response_id + mock_job.model_object_id = encrypted_response_id mock_job.created_by = "test-user" mock_job.id = "job-encrypted" mock_job.file_object = {"model": "gpt-5", "id": encrypted_response_id} @@ -556,7 +576,7 @@ class TestCheckResponsesCost: return_value=[mock_job] ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) mock_llm_router.aget_responses = AsyncMock( @@ -584,11 +604,7 @@ class TestCheckResponsesCost: mock_llm_router.aget_responses.call_args[1]["response_id"] == encoded_response_id ) - calls = ( - mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list - ) - assert len(calls) == 1 - assert calls[0][1]["where"]["id"]["in"] == ["job-encrypted"] + assert _completed_job_ids(mock_prisma_client) == ["job-encrypted"] @pytest.mark.asyncio async def test_response_id_without_model_id_uses_sdk( @@ -597,6 +613,7 @@ class TestCheckResponsesCost: """Ids that carry no deployment info can't be routed, so fall back to the SDK.""" mock_job = MagicMock() mock_job.unified_object_id = "resp_plain_upstream_id" + mock_job.model_object_id = "resp_plain_upstream_id" mock_job.created_by = "test-user" mock_job.id = "job-plain" mock_job.file_object = {"model": "gpt-5", "id": "resp_plain_upstream_id"} @@ -605,7 +622,7 @@ class TestCheckResponsesCost: return_value=[mock_job] ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) mock_llm_router.aget_responses = AsyncMock( side_effect=AssertionError("router cannot route an id without a model_id") @@ -624,7 +641,7 @@ class TestCheckResponsesCost: mock_sdk_aget.return_value = mock_response await check_responses_cost_instance.check_responses_cost() - mock_sdk_aget.assert_called_once() + assert mock_sdk_aget.await_count == 2 mock_llm_router.aget_responses.assert_not_called() @pytest.mark.asyncio @@ -646,6 +663,7 @@ class TestCheckResponsesCost: mock_job = MagicMock() mock_job.unified_object_id = encoded_response_id + mock_job.model_object_id = encoded_response_id mock_job.created_by = "test-user" mock_job.id = "job-missing-deployment" mock_job.file_object = {"model": "gpt-5", "id": encoded_response_id} @@ -654,7 +672,7 @@ class TestCheckResponsesCost: return_value=[mock_job] ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) mock_llm_router.get_deployment = MagicMock(return_value=None) mock_llm_router.aget_responses = AsyncMock( @@ -674,17 +692,16 @@ class TestCheckResponsesCost: mock_sdk_aget.return_value = mock_response await check_responses_cost_instance.check_responses_cost() - mock_llm_router.get_deployment.assert_called_once_with(model_id="deployment-deleted") + assert mock_llm_router.get_deployment.call_count == 2 + assert mock_llm_router.get_deployment.call_args_list == [ + call(model_id="deployment-deleted"), + call(model_id="deployment-deleted"), + ] mock_llm_router.aget_responses.assert_not_called() - mock_sdk_aget.assert_called_once() + assert mock_sdk_aget.await_count == 2 assert mock_sdk_aget.call_args[1]["response_id"] == encoded_response_id - calls = ( - mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list - ) - assert len(calls) == 1 - assert calls[0][1]["data"]["status"] == "completed" - assert calls[0][1]["where"]["id"]["in"] == ["job-missing-deployment"] + assert _completed_job_ids(mock_prisma_client) == ["job-missing-deployment"] @pytest.mark.asyncio async def test_check_responses_cost_with_incomplete_response( @@ -693,6 +710,7 @@ class TestCheckResponsesCost: """'incomplete' is terminal in the Responses API, so the row must not stay queued.""" mock_job = MagicMock() mock_job.unified_object_id = "resp_test_incomplete" + mock_job.model_object_id = "resp_test_incomplete" mock_job.created_by = "test-user" mock_job.id = "job-incomplete" mock_job.file_object = {"model": "gpt-5", "id": "resp_test_incomplete"} @@ -701,7 +719,7 @@ class TestCheckResponsesCost: return_value=[mock_job] ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) mock_response = ResponsesAPIResponse( @@ -717,12 +735,7 @@ class TestCheckResponsesCost: mock_aget.return_value = mock_response await check_responses_cost_instance.check_responses_cost() - calls = ( - mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list - ) - assert len(calls) == 1 - assert calls[0][1]["data"]["status"] == "completed" - assert calls[0][1]["where"]["id"]["in"] == ["job-incomplete"] + assert _completed_job_ids(mock_prisma_client) == ["job-incomplete"] @pytest.mark.asyncio async def test_check_responses_cost_no_model_in_file_object( @@ -731,7 +744,10 @@ class TestCheckResponsesCost: """When file_object has no 'model' key, model_name is None and metadata skips model fields.""" mock_job = MagicMock() mock_job.unified_object_id = "resp_test_no_model" + mock_job.model_object_id = "resp_test_no_model" mock_job.created_by = "test-user" + mock_job.team_id = None + mock_job.api_key = None mock_job.id = "job-no-model" mock_job.file_object = {} # no "model" key → model_name=None branch @@ -739,7 +755,7 @@ class TestCheckResponsesCost: return_value=[mock_job] ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) mock_response = MagicMock() @@ -753,6 +769,9 @@ class TestCheckResponsesCost: call_kwargs = mock_aget.call_args[1] assert "model" not in call_kwargs.get("litellm_metadata", {}) assert "model_group" not in call_kwargs.get("litellm_metadata", {}) + assert "user_api_key_team_id" not in call_kwargs["litellm_metadata"] + assert "user_api_key" not in call_kwargs["litellm_metadata"] + assert "user_api_key_hash" not in call_kwargs["litellm_metadata"] @pytest.mark.asyncio async def test_poll_stamps_internal_call_origin_so_the_read_is_billed( @@ -768,7 +787,10 @@ class TestCheckResponsesCost: mock_job = MagicMock() mock_job.unified_object_id = "resp_test_billed" + mock_job.model_object_id = "resp_test_billed" mock_job.created_by = "test-user" + mock_job.team_id = "team-billed" + mock_job.api_key = "sk-billed" mock_job.id = "job-billed" mock_job.file_object = {"model": "gpt-5", "id": "resp_test_billed"} @@ -776,7 +798,7 @@ class TestCheckResponsesCost: return_value=[mock_job] ) mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) mock_response = MagicMock() @@ -786,8 +808,736 @@ class TestCheckResponsesCost: mock_aget.return_value = mock_response await check_responses_cost_instance.check_responses_cost() - metadata = mock_aget.call_args[1]["litellm_metadata"] - foreground_read = {"background": False} - assert metadata[INTERNAL_CALL_ORIGIN_METADATA_KEY] == "background_response_cost_poll" - assert is_unbilled_non_inference_call("aget_responses", metadata, foreground_read) is False - assert is_unbilled_non_inference_call("aget_responses", None, foreground_read) is True + probe_metadata = mock_aget.call_args_list[0][1]["litellm_metadata"] + billing_metadata = mock_aget.call_args_list[1][1]["litellm_metadata"] + assert INTERNAL_CALL_ORIGIN_METADATA_KEY not in probe_metadata + assert billing_metadata[INTERNAL_CALL_ORIGIN_METADATA_KEY] == "background_response_cost_poll" + assert billing_metadata["user_api_key_team_id"] == "team-billed" + assert billing_metadata["user_api_key"] == "sk-billed" + assert billing_metadata["user_api_key_hash"] == "sk-billed" + assert is_unbilled_non_inference_call("aget_responses", probe_metadata) is True + assert is_unbilled_non_inference_call("aget_responses", billing_metadata) is False + assert is_unbilled_non_inference_call("aget_responses", None) is True + + @pytest.mark.asyncio + async def test_job_claimed_by_another_pod_is_never_read_or_completed( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + """Every pod and uvicorn worker polls the same table, and the read is what writes the + spend log, so losing the claim has to skip the read entirely or the job is billed twice.""" + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_claimed_elsewhere" + mock_job.model_object_id = _routed_response_id("resp_test_claimed_elsewhere") + mock_job.created_by = "test-user" + mock_job.id = "job-claimed-elsewhere" + mock_job.file_object = {"model": "gpt-5", "id": "resp_test_claimed_elsewhere"} + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + return_value=0 + ) + + await check_responses_cost_instance.check_responses_cost() + + mock_llm_router.aget_responses.assert_not_awaited() + assert _completion_calls(mock_prisma_client) == [] + assert _release_calls(mock_prisma_client) == [] + + claim_calls = _claim_calls(mock_prisma_client) + assert len(claim_calls) == 1 + claim_where = claim_calls[0].kwargs["where"] + assert claim_where["id"] == "job-claimed-elsewhere" + assert {"batch_processed": False} in claim_where["OR"] + + @pytest.mark.asyncio + async def test_claim_is_taken_back_from_a_pod_that_died_holding_it( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + """A pod that dies holding a claim strands the row forever, so the lease has to outlast a + live cycle and still fire well before stale expiry gives up on the row unbilled.""" + from litellm_enterprise.proxy.common_utils.check_responses_cost import ( + CLAIM_ABANDONED_AFTER_POLL_CYCLES, + ) + + from litellm.constants import PROXY_BATCH_POLLING_INTERVAL + + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_abandoned" + mock_job.model_object_id = _routed_response_id("resp_test_abandoned") + mock_job.created_by = "test-user" + mock_job.id = "job-abandoned" + mock_job.file_object = {"model": "gpt-5", "id": "resp_test_abandoned"} + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + return_value=1 + ) + + mock_response = ResponsesAPIResponse( + id="resp_abandoned", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage(input_tokens=100, output_tokens=50, total_tokens=150), + ) + + mock_llm_router.aget_responses = AsyncMock(return_value=mock_response) + + await check_responses_cost_instance.check_responses_cost() + + claim_where = _claim_calls(mock_prisma_client)[0].kwargs["where"] + abandoned_arm = next(arm for arm in claim_where["OR"] if "updated_at" in arm) + lease = timedelta( + seconds=CLAIM_ABANDONED_AFTER_POLL_CYCLES * PROXY_BATCH_POLLING_INTERVAL + ) + untouched_for = datetime.now(timezone.utc) - abandoned_arm["updated_at"]["lt"] + assert lease <= untouched_for < lease + timedelta(seconds=30) + assert lease > timedelta(seconds=PROXY_BATCH_POLLING_INTERVAL) + + @pytest.mark.asyncio + async def test_claim_is_taken_before_the_billing_read_and_kept_on_a_terminal_status( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + """The read prices the job, so the claim has to be taken before it, and keeping the claim + afterwards is what stops a second pod reading and billing the same row again.""" + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_ordering" + mock_job.model_object_id = _routed_response_id("resp_test_ordering") + mock_job.created_by = "test-user" + mock_job.id = "job-ordering" + mock_job.file_object = {"model": "gpt-5", "id": "resp_test_ordering"} + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + + writes_and_reads = [] + + async def record_update_many(**kwargs): + writes_and_reads.append(kwargs) + return 1 + + async def record_read(**kwargs): + writes_and_reads.append("provider_read") + return ResponsesAPIResponse( + id="resp_ordering", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage( + input_tokens=100, output_tokens=50, total_tokens=150 + ), + ) + + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + side_effect=record_update_many + ) + + mock_llm_router.aget_responses = AsyncMock(side_effect=record_read) + + await check_responses_cost_instance.check_responses_cost() + + assert len(writes_and_reads) == 4 + assert writes_and_reads[0]["data"] == {"batch_processed": True} + assert writes_and_reads[1] == "provider_read" + assert writes_and_reads[2]["data"] == {"status": "completed"} + assert writes_and_reads[2]["where"] == { + "id": "job-ordering", + "status": {"in": ["queued", "in_progress"]}, + } + assert writes_and_reads[3] == "provider_read" + + @pytest.mark.asyncio + async def test_probe_read_is_free_and_only_the_post_cas_read_is_billed( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY + from litellm.types.utils import BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN + + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_probe_billing" + mock_job.model_object_id = _routed_response_id("resp_test_probe_billing") + mock_job.created_by = "test-user" + mock_job.id = "job-probe-billing" + mock_job.file_object = {"model": "gpt-5", "id": "resp_test_probe_billing"} + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + + response = ResponsesAPIResponse( + id="resp_probe_billing", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage(input_tokens=10, output_tokens=5, total_tokens=15), + ) + call_order = [] + + async def record_update_many(**kwargs): + call_order.append(("update", kwargs)) + return 1 + + async def record_read(**kwargs): + call_order.append(("read", kwargs)) + return response + + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + side_effect=record_update_many + ) + mock_llm_router.aget_responses = AsyncMock(side_effect=record_read) + + await check_responses_cost_instance.check_responses_cost() + + assert mock_llm_router.aget_responses.await_count == 2 + assert [kind for kind, _ in call_order] == ["update", "read", "update", "read"] + probe_metadata = call_order[1][1]["litellm_metadata"] + billing_metadata = call_order[3][1]["litellm_metadata"] + assert INTERNAL_CALL_ORIGIN_METADATA_KEY not in probe_metadata + assert billing_metadata[INTERNAL_CALL_ORIGIN_METADATA_KEY] == ( + BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN + ) + assert call_order[2][1]["where"] == { + "id": "job-probe-billing", + "status": {"in": ["queued", "in_progress"]}, + } + + @pytest.mark.asyncio + async def test_row_already_finalized_by_another_poller_is_not_billed( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY + + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_finalized" + mock_job.model_object_id = _routed_response_id("resp_test_finalized") + mock_job.created_by = "test-user" + mock_job.id = "job-finalized" + mock_job.file_object = {"model": "gpt-5", "id": "resp_test_finalized"} + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + side_effect=[1, 0] + ) + mock_llm_router.aget_responses = AsyncMock( + return_value=ResponsesAPIResponse( + id="resp_finalized", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + ) + + await check_responses_cost_instance.check_responses_cost() + + assert mock_llm_router.aget_responses.await_count == 1 + metadata = mock_llm_router.aget_responses.call_args.kwargs["litellm_metadata"] + assert INTERNAL_CALL_ORIGIN_METADATA_KEY not in metadata + completion_calls = _completion_calls(mock_prisma_client) + assert len(completion_calls) == 1 + assert completion_calls[0].kwargs["where"]["status"] == { + "in": ["queued", "in_progress"] + } + assert _release_calls(mock_prisma_client) == [] + + @pytest.mark.asyncio + async def test_billing_read_includes_managed_row_attribution( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY + from litellm.types.utils import BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN + + attributed_job = MagicMock() + attributed_job.unified_object_id = "resp_attributed" + attributed_job.model_object_id = _routed_response_id("resp_attributed") + attributed_job.created_by = "test-user" + attributed_job.org_id = "org-1" + attributed_job.request_tags = ["tag-a", "tag-b"] + attributed_job.id = "job-attributed" + attributed_job.file_object = {"model": "gpt-5", "id": "resp_attributed"} + + unattributed_job = MagicMock() + unattributed_job.unified_object_id = "resp_unattributed" + unattributed_job.model_object_id = _routed_response_id("resp_unattributed") + unattributed_job.created_by = "test-user" + unattributed_job.org_id = None + unattributed_job.request_tags = [] + unattributed_job.id = "job-unattributed" + unattributed_job.file_object = {"model": "gpt-5", "id": "resp_unattributed"} + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[attributed_job, unattributed_job] + ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) + mock_llm_router.aget_responses = AsyncMock( + side_effect=[ + ResponsesAPIResponse( + id="resp_attributed", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + ] + * 2 + + [ + ResponsesAPIResponse( + id="resp_unattributed", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + ] + * 2 + ) + + await check_responses_cost_instance.check_responses_cost() + + billing_metadata = [ + call.kwargs["litellm_metadata"] + for call in mock_llm_router.aget_responses.await_args_list + if call.kwargs["litellm_metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) + == BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN + ] + assert billing_metadata[0]["user_api_key_org_id"] == "org-1" + assert billing_metadata[0]["tags"] == ["tag-a", "tag-b"] + assert "user_api_key_org_id" not in billing_metadata[1] + assert "tags" not in billing_metadata[1] + + @pytest.mark.asyncio + async def test_non_terminal_probe_does_not_finalize_or_bill( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_non_terminal" + mock_job.model_object_id = _routed_response_id("resp_test_non_terminal") + mock_job.created_by = "test-user" + mock_job.id = "job-non-terminal" + mock_job.file_object = {"model": "gpt-5", "id": "resp_test_non_terminal"} + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + return_value=1 + ) + mock_llm_router.aget_responses = AsyncMock( + return_value=ResponsesAPIResponse( + id="resp_non_terminal", + object="response", + status="queued", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + ) + + await check_responses_cost_instance.check_responses_cost() + + assert mock_llm_router.aget_responses.await_count == 1 + assert _completion_calls(mock_prisma_client) == [] + assert len(_release_calls(mock_prisma_client)) == 1 + + @pytest.mark.asyncio + async def test_billing_read_failure_after_cas_does_not_reopen_the_row( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_billing_failure" + mock_job.model_object_id = _routed_response_id("resp_test_billing_failure") + mock_job.created_by = "test-user" + mock_job.id = "job-billing-failure" + mock_job.file_object = {"model": "gpt-5", "id": "resp_test_billing_failure"} + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + return_value=1 + ) + mock_llm_router.aget_responses = AsyncMock( + side_effect=[ + ResponsesAPIResponse( + id="resp_billing_failure", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ), + Exception("boom"), + ] + ) + + await check_responses_cost_instance.check_responses_cost() + + assert mock_llm_router.aget_responses.await_count == 2 + assert len(_completion_calls(mock_prisma_client)) == 1 + assert _release_calls(mock_prisma_client) == [] + assert all( + call.kwargs["data"] not in ({"batch_processed": False}, {"status": "queued"}, {"status": "in_progress"}) + for call in mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize("provider_status", ["queued", "in_progress"]) + async def test_non_terminal_status_releases_the_claim( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router, provider_status + ): + """A response the provider has not finished yet has no spend to record, so its row must go + back to batch_processed=False; holding the claim retires it before it is ever billed.""" + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_still_running" + mock_job.model_object_id = _routed_response_id("resp_test_still_running") + mock_job.created_by = "test-user" + mock_job.id = "job-still-running" + mock_job.file_object = {"model": "gpt-5", "id": "resp_test_still_running"} + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + return_value=1 + ) + + mock_response = ResponsesAPIResponse( + id="resp_still_running", + object="response", + status=provider_status, + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + + mock_llm_router.aget_responses = AsyncMock(return_value=mock_response) + + await check_responses_cost_instance.check_responses_cost() + + assert _completion_calls(mock_prisma_client) == [] + release_calls = _release_calls(mock_prisma_client) + assert len(release_calls) == 1 + assert release_calls[0].kwargs["where"] == { + "id": "job-still-running", + "batch_processed": True, + } + + @pytest.mark.asyncio + async def test_failed_provider_read_releases_the_claim( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + """A read that raised billed nothing, so the claim has to be handed back or the row is + retired unbilled and no later poll cycle ever retries it.""" + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_read_error" + mock_job.model_object_id = _routed_response_id("resp_test_read_error") + mock_job.created_by = "test-user" + mock_job.id = "job-read-error" + mock_job.file_object = {"model": "gpt-5", "id": "resp_test_read_error"} + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + return_value=1 + ) + + mock_llm_router.aget_responses = AsyncMock(side_effect=Exception("Provider error")) + + await check_responses_cost_instance.check_responses_cost() + + assert _completion_calls(mock_prisma_client) == [] + release_calls = _release_calls(mock_prisma_client) + assert len(release_calls) == 1 + assert release_calls[0].kwargs["where"] == { + "id": "job-read-error", + "batch_processed": True, + } + + @pytest.mark.asyncio + async def test_a_job_claimed_elsewhere_does_not_block_the_next_job( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + """Losing one row to another pod must skip only that row: the rest of the poll page still + has to be read and billed in the same cycle.""" + mock_job1 = MagicMock() + mock_job1.unified_object_id = "resp_test_first" + mock_job1.model_object_id = _routed_response_id("resp_test_first") + mock_job1.created_by = "user1" + mock_job1.id = "job-first" + mock_job1.file_object = {"model": "gpt-5", "id": "resp_test_first"} + + mock_job2 = MagicMock() + mock_job2.unified_object_id = "resp_test_second" + mock_job2.model_object_id = _routed_response_id("resp_test_second") + mock_job2.created_by = "user2" + mock_job2.id = "job-second" + mock_job2.file_object = {"model": "gpt-5", "id": "resp_test_second"} + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job1, mock_job2] + ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + side_effect=[0, 1, 1] + ) + + mock_response = ResponsesAPIResponse( + id="resp_second", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage(input_tokens=100, output_tokens=50, total_tokens=150), + ) + + mock_llm_router.aget_responses = AsyncMock(return_value=mock_response) + + await check_responses_cost_instance.check_responses_cost() + + assert mock_llm_router.aget_responses.await_count == 2 + assert mock_llm_router.aget_responses.await_args.kwargs["response_id"] == _routed_response_id( + "resp_test_second" + ) + + assert _completed_job_ids(mock_prisma_client) == ["job-second"] + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "db_error_message", + [ + "column LiteLLM_ManagedObjectTable.batch_processed does not exist", + "Unknown column in where clause", + "The column P2022 does not exist in the current database", + ], + ) + async def test_claim_fails_open_on_a_schema_without_the_claim_column( + self, check_responses_cost_instance, mock_prisma_client, db_error_message + ): + """A deployment that never ran the batch_processed migration cannot claim anything, so it + keeps the pre-claim behavior of billing rather than silently billing nothing.""" + mock_job = MagicMock() + mock_job.id = "job-old-schema" + + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + side_effect=Exception(db_error_message) + ) + + assert ( + await check_responses_cost_instance._claim_job_for_costing(mock_job) is True + ) + + @pytest.mark.asyncio + async def test_claim_is_lost_when_the_database_fails_for_any_other_reason( + self, check_responses_cost_instance, mock_prisma_client + ): + """A dropped connection is no proof the row is free, so the read that would bill it is + not allowed to run.""" + mock_job = MagicMock() + mock_job.id = "job-db-down" + + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + side_effect=Exception("connection to server was lost") + ) + + assert ( + await check_responses_cost_instance._claim_job_for_costing(mock_job) is False + ) + + @pytest.mark.asyncio + async def test_old_schema_without_the_claim_column_still_bills_and_completes( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + """End to end on a pre-migration schema: the claim write fails, the response is still read + (which is what bills it) and the row is still marked completed.""" + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_old_schema" + mock_job.model_object_id = _routed_response_id("resp_test_old_schema") + mock_job.created_by = "test-user" + mock_job.id = "job-old-schema" + mock_job.file_object = {"model": "gpt-5", "id": "resp_test_old_schema"} + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + + async def reject_batch_processed_writes(**kwargs): + if "batch_processed" in kwargs["data"]: + raise Exception( + 'column "batch_processed" of relation ' + '"LiteLLM_ManagedObjectTable" does not exist' + ) + return 1 + + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + side_effect=reject_batch_processed_writes + ) + + mock_response = ResponsesAPIResponse( + id="resp_old_schema", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage(input_tokens=100, output_tokens=50, total_tokens=150), + ) + + mock_llm_router.aget_responses = AsyncMock(return_value=mock_response) + + await check_responses_cost_instance.check_responses_cost() + + assert mock_llm_router.aget_responses.await_count == 2 + assert _completed_job_ids(mock_prisma_client) == ["job-old-schema"] + + @pytest.mark.asyncio + async def test_a_failed_persist_does_not_abort_the_rest_of_the_poll_cycle( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + """One row's completion write failing must not take the whole cycle down with it.""" + mock_job1 = MagicMock() + mock_job1.unified_object_id = "resp_test_persist_fails" + mock_job1.model_object_id = _routed_response_id("resp_test_persist_fails") + mock_job1.created_by = "user1" + mock_job1.id = "job-persist-fails" + mock_job1.file_object = {"model": "gpt-5", "id": "resp_test_persist_fails"} + + mock_job2 = MagicMock() + mock_job2.unified_object_id = "resp_test_persist_works" + mock_job2.model_object_id = _routed_response_id("resp_test_persist_works") + mock_job2.created_by = "user2" + mock_job2.id = "job-persist-works" + mock_job2.file_object = {"model": "gpt-5", "id": "resp_test_persist_works"} + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job1, mock_job2] + ) + async def fail_first_completion_write(**kwargs): + if ( + kwargs["data"] == {"status": "completed"} + and kwargs["where"]["id"] == "job-persist-fails" + ): + raise Exception("deadlock detected") + return 1 + + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + side_effect=fail_first_completion_write + ) + + mock_response = ResponsesAPIResponse( + id="resp_persisted", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage(input_tokens=100, output_tokens=50, total_tokens=150), + ) + + mock_llm_router.aget_responses = AsyncMock(return_value=mock_response) + + await check_responses_cost_instance.check_responses_cost() + + assert mock_llm_router.aget_responses.await_count == 3 + assert _completed_job_ids(mock_prisma_client) == [ + "job-persist-fails", + "job-persist-works", + ] + + @pytest.mark.asyncio + async def test_poller_fetches_the_provider_id_from_model_object_id( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router, monkeypatch + ): + """The row's provider id drives the fetch, not the nonce-encrypted advertised id. + + A background create advertises a freshly encrypted id per call, so unified_object_id + is no handle on the generation. + """ + from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper + from litellm.types.utils import SpecialEnums + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-key-for-response-ids") + + provider_response_id = _routed_response_id("resp_upstream_stable") + stale_advertised_id = "resp_" + str( + encrypt_value_helper( + value=SpecialEnums.LITELLM_MANAGED_RESPONSE_API_RESPONSE_ID_COMPLETE_STR.value.format( + "resp_a_previous_encoding", "test-user", "test-team" + ) + ) + ) + + mock_job = MagicMock() + mock_job.unified_object_id = stale_advertised_id + mock_job.model_object_id = provider_response_id + mock_job.created_by = "test-user" + mock_job.id = "job-provider-id" + mock_job.file_object = {"model": "gpt-5", "id": stale_advertised_id} + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[mock_job]) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) + mock_llm_router.aget_responses = AsyncMock( + return_value=ResponsesAPIResponse( + id=provider_response_id, + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage(input_tokens=10, output_tokens=5, total_tokens=15), + ) + ) + + await check_responses_cost_instance.check_responses_cost() + + assert mock_llm_router.aget_responses.call_args[1]["response_id"] == provider_response_id + assert _completed_job_ids(mock_prisma_client) == ["job-provider-id"] + + @pytest.mark.asyncio + async def test_poller_still_reads_rows_written_before_the_provider_id_was_stored( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router, monkeypatch + ): + """Rows created earlier carry the encrypted advertised id in both columns.""" + from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper + from litellm.types.utils import SpecialEnums + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-key-for-response-ids") + + provider_response_id = _routed_response_id("resp_legacy_upstream") + legacy_id = "resp_" + str( + encrypt_value_helper( + value=SpecialEnums.LITELLM_MANAGED_RESPONSE_API_RESPONSE_ID_COMPLETE_STR.value.format( + provider_response_id, "test-user", "test-team" + ) + ) + ) + + mock_job = MagicMock() + mock_job.unified_object_id = legacy_id + mock_job.model_object_id = legacy_id + mock_job.created_by = "test-user" + mock_job.id = "job-legacy" + mock_job.file_object = {"model": "gpt-5", "id": legacy_id} + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[mock_job]) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) + mock_llm_router.aget_responses = AsyncMock( + return_value=ResponsesAPIResponse( + id=provider_response_id, + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage(input_tokens=10, output_tokens=5, total_tokens=15), + ) + ) + + await check_responses_cost_instance.check_responses_cost() + + assert mock_llm_router.aget_responses.call_args[1]["response_id"] == provider_response_id + assert _completed_job_ids(mock_prisma_client) == ["job-legacy"] diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_metrics.py b/tests/test_litellm/integrations/otel/test_otel_v2_metrics.py index 016dbcd824b..32f1bd8ceda 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_metrics.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_metrics.py @@ -214,10 +214,8 @@ def test_response_read_does_not_replay_the_generation_usage(): assert RESPONSE_DURATION in metrics -def test_background_response_read_still_records_usage(): - """A background=true create returns no usage, so its completed read is the only - place the generation's tokens are ever seen. Skipping it would lose them - entirely rather than deduplicate them.""" +def test_background_response_read_does_not_record_usage(): + """The enterprise cost poller owns usage metrics for completed background responses.""" reader = InMemoryMetricReader() logger = _logger(reader, enable_metrics=True) kwargs, response_obj, start, end = _build_call(call_type="aget_responses") @@ -225,10 +223,8 @@ def test_background_response_read_still_records_usage(): asyncio.run(logger.async_log_success_event(kwargs, response_obj, start, end)) metrics = _metrics_by_name(reader) - by_type = {dp.attributes[TOKEN_TYPE]: dp for dp in metrics[TOKEN_USAGE]} - assert by_type["input"].sum == PROMPT_TOKENS - assert by_type["output"].sum == COMPLETION_TOKENS - assert TIME_PER_OUTPUT_TOKEN in metrics + assert TOKEN_USAGE not in metrics + assert TIME_PER_OUTPUT_TOKEN not in metrics def test_metrics_disabled_records_nothing(): diff --git a/tests/test_litellm/integrations/test_opentelemetry.py b/tests/test_litellm/integrations/test_opentelemetry.py index 974961f2eb5..0b51a8d2dbd 100644 --- a/tests/test_litellm/integrations/test_opentelemetry.py +++ b/tests/test_litellm/integrations/test_opentelemetry.py @@ -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 ) diff --git a/tests/test_litellm/integrations/test_responses_background_cost.py b/tests/test_litellm/integrations/test_responses_background_cost.py index 0d4218f2137..8e5e72bf032 100644 --- a/tests/test_litellm/integrations/test_responses_background_cost.py +++ b/tests/test_litellm/integrations/test_responses_background_cost.py @@ -81,6 +81,7 @@ class TestResponsesBackgroundCostTracking: model_object_id=response.id, file_purpose="response", user_api_key_dict=user_api_key_dict, + persist_attribution=True, ) # Verify store_unified_object_id was called @@ -92,6 +93,7 @@ class TestResponsesBackgroundCostTracking: assert call_args[1]["model_object_id"] == response.id assert call_args[1]["file_purpose"] == "response" assert call_args[1]["user_api_key_dict"] == user_api_key_dict + assert call_args[1]["persist_attribution"] is True @pytest.mark.asyncio async def test_no_storage_for_non_background_requests( @@ -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 diff --git a/tests/test_litellm/litellm_core_utils/test_internal_call_metadata.py b/tests/test_litellm/litellm_core_utils/test_internal_call_metadata.py index 73923dc75a5..e69e5313aae 100644 --- a/tests/test_litellm/litellm_core_utils/test_internal_call_metadata.py +++ b/tests/test_litellm/litellm_core_utils/test_internal_call_metadata.py @@ -3,9 +3,10 @@ from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY from litellm.litellm_core_utils.internal_call_metadata import ( forwarded_internal_call_metadata, + is_unbilled_non_inference_call, sanitized_forwardable_call_metadata, ) -from litellm.types.utils import SHADOW_EVAL_ROUTER_CALL_ORIGIN +from litellm.types.utils import BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN, SHADOW_EVAL_ROUTER_CALL_ORIGIN PARENT = { "user_api_key": "sk-hash", @@ -49,6 +50,17 @@ def test_sanitized_forwardable_metadata_keeps_only_identity_and_always_stamps(): } +def test_background_response_reads_are_free_but_cost_poller_reads_are_billed(): + assert is_unbilled_non_inference_call("aget_responses", {}) is True + assert ( + is_unbilled_non_inference_call( + "aget_responses", + {INTERNAL_CALL_ORIGIN_METADATA_KEY: BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN}, + ) + is False + ) + + class TestSubCallMetadataSanitization: """The proxy cost callback must not be able to recover the parent budget reservation from sub-call metadata, in either of the shapes it knows how to read.""" diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index 3d5c38c3acd..6528785e318 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -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 diff --git a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py index 4153bf7d7ee..7289cbab97e 100644 --- a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py @@ -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 diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py index f471e3f8fbb..4362f425fb8 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py @@ -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(): diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 5b9cd761dda..8060086919b 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -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): diff --git a/tests/test_litellm/test_responses_id_security.py b/tests/test_litellm/test_responses_id_security.py index a6081670172..44504194686 100644 --- a/tests/test_litellm/test_responses_id_security.py +++ b/tests/test_litellm/test_responses_id_security.py @@ -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"