From c98c0dc1ba9f097ecb053f09a935ab9e9aab6c89 Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Fri, 19 Jun 2026 12:03:15 -0700 Subject: [PATCH] fix(proxy): record partial spend on the failure row for interrupted streams (#30788) A streaming request that breaks mid-flight, for example on a mid-stream read timeout, still bills the provider for the chunks already delivered, yet the proxy recorded that interrupted request as a zero-spend failure. An earlier revision logged the recovered partial usage through the success path, which mislabeled a failed request as a success and produced a misleading spend row This recovers the partial usage where the failure is actually logged. The streaming handler assembles the usage from the chunks seen so far and stashes it, with its cost, on the logging object before firing the failure handlers. The proxy failure hook lifts that usage and cost onto request_data before the non-serialisable logging object is popped, and the spend-log writer records the real partial spend on the failure row instead of a hardcoded zero; get_logging_payload honors the recovered usage for the token columns and _failure_handler_helper_fn preserves the recovered cost so the non-DB failure loggers stay consistent A request that recovers via a successful fallback is unaffected: the failure hook only fires when the whole request fails, so the fallback's combined-usage success row stays the single source of truth and there is no double counting Resolves LIT-3825 Co-authored-by: veria-ai[bot] <224490171+veria-ai[bot]@users.noreply.github.com> (cherry picked from commit 4847fa5dd5991496a071d235781e07d39857b0f7) --- litellm/litellm_core_utils/litellm_logging.py | 7 +- .../litellm_core_utils/streaming_handler.py | 29 + .../proxy/hooks/proxy_track_cost_callback.py | 1139 +++++++++-------- .../spend_tracking/spend_tracking_utils.py | 7 + litellm/proxy/utils.py | 14 + .../test_litellm_logging.py | 43 + .../test_streaming_handler.py | 76 ++ .../hooks/test_proxy_track_cost_callback.py | 37 + .../test_spend_tracking_utils.py | 47 + tests/test_litellm/proxy/test_proxy_utils.py | 49 + tests/test_litellm/test_router.py | 145 +++ 11 files changed, 1028 insertions(+), 565 deletions(-) diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index e12a8365eb5..a24fd070a26 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -2944,7 +2944,12 @@ class Logging(LiteLLMLoggingBaseClass): ) self.model_call_details["end_time"] = end_time self.model_call_details.setdefault("original_response", None) - self.model_call_details["response_cost"] = 0 + # A stream interrupted mid-flight still billed the provider for the + # chunks already delivered; the router stashes that recovered usage as + # ``combined_usage_object`` and pre-computes its cost, so preserve it + # here instead of zeroing the spend on an otherwise-failed request. + if self.model_call_details.get("combined_usage_object") is None: + self.model_call_details["response_cost"] = 0 if hasattr(exception, "headers") and isinstance(exception.headers, dict): self.model_call_details.setdefault("litellm_params", {}) diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 29c0d0629e8..3d04b183176 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -2233,6 +2233,7 @@ class CustomStreamWrapper: litellm.request_timeout ) if self.logging_obj is not None: + self._record_partial_usage_for_failure() ## LOGGING threading.Thread( target=self.logging_obj.failure_handler, @@ -2246,6 +2247,7 @@ class CustomStreamWrapper: except Exception as e: traceback_exception = traceback.format_exc() if self.logging_obj is not None: + self._record_partial_usage_for_failure() ## LOGGING threading.Thread( target=self.logging_obj.failure_handler, @@ -2257,6 +2259,33 @@ class CustomStreamWrapper: ) self._handle_stream_fallback_error(e) + def _record_partial_usage_for_failure(self) -> None: + """ + A stream that breaks mid-flight still billed the provider for the chunks + already delivered. Recover that partial usage from the chunks seen so + far and stash it, with its cost, on the logging object so the failure + handler records the real partial spend instead of zero. A request that + later recovers via a router fallback overwrites this with the combined + success log on the same request id, so this never double counts. + """ + if self.logging_obj is None or not self.chunks: + return + try: + partial_response = litellm.stream_chunk_builder(chunks=self.chunks) + usage = cast(Optional[Usage], getattr(partial_response, "usage", None)) + if usage is None: + return + self.logging_obj.model_call_details["combined_usage_object"] = usage + self.logging_obj.model_call_details["response_cost"] = ( + self.logging_obj._response_cost_calculator(result=partial_response) + or 0.0 + ) + except Exception as recover_error: + verbose_logger.debug( + "could not recover partial usage for interrupted stream: %s", + recover_error, + ) + def _handle_stream_fallback_error(self, e: Exception) -> "NoReturn": """ Common error handling for both __next__ and __anext__. diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 3688f25ac44..9d2b8572774 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -1,564 +1,575 @@ -import asyncio -import traceback -from datetime import datetime -from typing import Any, List, Optional, Union, cast - -import litellm -from litellm._logging import verbose_proxy_logger -from litellm.integrations.custom_logger import CustomLogger -from litellm.litellm_core_utils.core_helpers import ( - _get_parent_otel_span_from_kwargs, - get_litellm_metadata_from_kwargs, -) -from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup -from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.auth.auth_checks import ( - get_key_object, - get_team_object, - log_db_metrics, -) -from litellm.proxy.auth.route_checks import RouteChecks -from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup -from litellm.proxy.spend_tracking.spend_log_error_logger import ( - should_suppress_spend_log_tracebacks, - spend_log_error, -) -from litellm.proxy.spend_tracking.spend_tracking_utils import ( - _sanitize_error_information_for_spend_logs, -) -from litellm.proxy.utils import ProxyUpdateSpend -from litellm.types.utils import ( - StandardLoggingPayload, - StandardLoggingPayloadErrorInformation, -) -from litellm.utils import get_end_user_id_for_cost_tracking - - -class _ProxyDBLogger(CustomLogger): - async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): - await self._PROXY_track_cost_callback( - kwargs, response_obj, start_time, end_time - ) - - async def async_post_call_failure_hook( - self, - request_data: dict, - original_exception: Exception, - user_api_key_dict: UserAPIKeyAuth, - traceback_str: Optional[str] = None, - ): - try: - await _release_budget_reservation( - budget_reservation=user_api_key_dict.budget_reservation - ) - except Exception: - verbose_proxy_logger.exception( - "Failed to release budget reservation during failure handling" - ) - try: - await _invalidate_budget_reservation_counters( - budget_reservation=user_api_key_dict.budget_reservation - ) - if user_api_key_dict.budget_reservation is not None: - user_api_key_dict.budget_reservation["finalized"] = True - except Exception: - verbose_proxy_logger.exception( - "Failed to invalidate budget reservation counters after failure release failed" - ) - - request_route = user_api_key_dict.request_route - if _ProxyDBLogger._should_track_errors_in_db() is False: - return - elif request_route is not None and not ( - RouteChecks.is_llm_api_route(route=request_route) - or RouteChecks.is_info_route(route=request_route) - ): - return - - from litellm.proxy.proxy_server import proxy_logging_obj - - _metadata = dict( - LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key( - user_api_key_dict=user_api_key_dict - ) - ) - _metadata["user_api_key"] = user_api_key_dict.api_key - _metadata["status"] = "failure" - _error_information = StandardLoggingPayloadSetup.get_error_information( - original_exception=original_exception, - traceback_str=traceback_str, - ) - if should_suppress_spend_log_tracebacks(): - # Drop the traceback key entirely so the per-row Metadata pane in - # the UI (which renders the JSON blob verbatim) doesn't show a - # noisy ``"traceback": ""`` line. Downstream consumers all use - # ``.get("traceback")`` / truthy checks, and the TypedDict marks - # the field as optional, so omitting is type-safe. - _error_information.pop("traceback", None) - # Strip echoed request input + apply DB-size cap before storing in - # the spend-log metadata column (LIT-2992). Result is never None - # here because the input above is constructed non-None. - _error_information = cast( - StandardLoggingPayloadErrorInformation, - _sanitize_error_information_for_spend_logs(_error_information), - ) - _metadata["error_information"] = _error_information - - _metadata = await _ProxyDBLogger._enrich_failure_metadata_with_key_info( - metadata=_metadata, - ) - - existing_metadata: dict = request_data.get("metadata", None) or {} - existing_metadata.update(_metadata) - - if "litellm_params" not in request_data: - request_data["litellm_params"] = {} - - existing_litellm_params = request_data.get("litellm_params", {}) - existing_litellm_metadata = existing_litellm_params.get("metadata", {}) or {} - - # Preserve tags from existing metadata - if existing_litellm_metadata.get("tags"): - existing_metadata["tags"] = existing_litellm_metadata.get("tags") - - request_data["litellm_params"]["proxy_server_request"] = ( - request_data.get("proxy_server_request") - or existing_litellm_params.get("proxy_server_request") - or {} - ) - request_data["litellm_params"]["metadata"] = existing_metadata - - # Preserve model name and custom_llm_provider - if "model" not in request_data: - request_data["model"] = existing_litellm_params.get( - "model" - ) or request_data.get("model", "") - if "custom_llm_provider" not in request_data: - request_data["custom_llm_provider"] = existing_litellm_params.get( - "custom_llm_provider" - ) or request_data.get("custom_llm_provider", "") - - # Propagate standard_logging_object and litellm_trace_id from the - # Logging instance so that _get_session_id_for_spend_log uses the same - # trace_id that Langfuse received (via async_failure_handler). - # Without this, the DB session_id would be a random UUID that doesn't - # match the Langfuse trace_id, making failed requests unsearchable. - _litellm_logging_obj = request_data.get("litellm_logging_obj") - if _litellm_logging_obj is not None: - if not request_data.get("standard_logging_object"): - request_data["standard_logging_object"] = getattr( - _litellm_logging_obj, "model_call_details", {} - ).get("standard_logging_object") - if request_data.get("litellm_trace_id") is None: - request_data["litellm_trace_id"] = getattr( - _litellm_logging_obj, "litellm_trace_id", None - ) - - # Use the actual request start time from the logging object so that - # failed requests record the real duration instead of 0. - actual_start_time = datetime.now() - if _litellm_logging_obj is not None: - obj_start = getattr(_litellm_logging_obj, "start_time", None) - if obj_start is not None: - actual_start_time = obj_start - - await proxy_logging_obj.db_spend_update_writer.update_database( - token=user_api_key_dict.api_key, - response_cost=0.0, - user_id=user_api_key_dict.user_id, - end_user_id=user_api_key_dict.end_user_id, - team_id=user_api_key_dict.team_id, - kwargs=request_data, - completion_response=original_exception, - start_time=actual_start_time, - end_time=datetime.now(), - org_id=user_api_key_dict.org_id, - ) - - @log_db_metrics - async def _PROXY_track_cost_callback( - self, - kwargs, # kwargs to completion - completion_response: Optional[ - Union[litellm.ModelResponse, Any] - ], # response from completion - start_time=None, - end_time=None, # start/end time for completion - ): - from litellm.proxy.proxy_server import ( - increment_spend_counters, - proxy_logging_obj, - update_cache, - ) - - verbose_proxy_logger.debug("INSIDE _PROXY_track_cost_callback") - try: - verbose_proxy_logger.debug( - f"kwargs stream: {kwargs.get('stream', None)} + complete streaming response: {kwargs.get('complete_streaming_response', None)}" - ) - parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs=kwargs) - litellm_params = kwargs.get("litellm_params", {}) or {} - end_user_id = get_end_user_id_for_cost_tracking(litellm_params) - metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs) - budget_reservation = _get_budget_reservation_from_metadata( - metadata=metadata - ) - user_id = cast(Optional[str], metadata.get("user_api_key_user_id", None)) - team_id = cast(Optional[str], metadata.get("user_api_key_team_id", None)) - org_id = cast(Optional[str], metadata.get("user_api_key_org_id", None)) - key_alias = cast(Optional[str], metadata.get("user_api_key_alias", None)) - end_user_max_budget = metadata.get("user_api_end_user_max_budget", None) - sl_object: Optional[StandardLoggingPayload] = kwargs.get( - "standard_logging_object", None - ) - response_cost = ( - sl_object.get("response_cost", None) - if sl_object is not None - else kwargs.get("response_cost", None) - ) - tags = _get_request_tags_for_cost_tracking( - sl_object=sl_object, - metadata=metadata, - ) - - if response_cost is not None: - user_api_key = metadata.get("user_api_key", None) - if kwargs.get("cache_hit", False) is True: - response_cost = 0.0 - verbose_proxy_logger.debug( - f"Cache Hit: response_cost {response_cost}, for user_id {user_id}" - ) - - verbose_proxy_logger.debug( - f"user_api_key {user_api_key}, user_id {user_id}, team_id {team_id}, end_user_id {end_user_id}" - ) - if _should_track_cost_callback( - user_api_key=user_api_key, - user_id=user_id, - team_id=team_id, - end_user_id=end_user_id, - ): - ## UPDATE DATABASE - await _update_database_and_spend_counters( - proxy_logging_obj=proxy_logging_obj, - increment_spend_counters=increment_spend_counters, - user_api_key=user_api_key, - user_id=user_id, - end_user_id=end_user_id, - team_id=team_id, - org_id=org_id, - kwargs=kwargs, - completion_response=completion_response, - start_time=start_time, - end_time=end_time, - response_cost=response_cost, - budget_reservation=budget_reservation, - request_tags=tags, - ) - - # update cache (fire-and-forget for backward compat: - # cached object fields, soft budget alerts, etc.) - asyncio.create_task( - update_cache( - token=user_api_key, - user_id=user_id, - end_user_id=end_user_id, - response_cost=response_cost, - team_id=team_id, - parent_otel_span=parent_otel_span, - tags=tags, - ) - ) - - await proxy_logging_obj.slack_alerting_instance.customer_spend_alert( - token=user_api_key, - key_alias=key_alias, - end_user_id=end_user_id, - response_cost=response_cost, - max_budget=end_user_max_budget, - ) - elif budget_reservation is not None: - await _release_budget_reservation( - budget_reservation=budget_reservation - ) - else: - await _release_budget_reservation(budget_reservation=budget_reservation) - # Non-model call types (health checks, afile_delete) have no model or standard_logging_object. - # Use .get() for "stream" to avoid KeyError on health checks. - if sl_object is None and not kwargs.get("model"): - verbose_proxy_logger.warning( - "Cost tracking - skipping, no standard_logging_object and no model for call_type=%s", - kwargs.get("call_type", "unknown"), - ) - return - if kwargs.get("stream") is not True or ( - kwargs.get("stream") is True - and "complete_streaming_response" in kwargs - ): - if sl_object is not None: - cost_tracking_failure_debug_info: Union[dict, str] = ( - sl_object["response_cost_failure_debug_info"] # type: ignore - or "response_cost_failure_debug_info is None in standard_logging_object" - ) - else: - cost_tracking_failure_debug_info = ( - "standard_logging_object not found" - ) - model = kwargs.get("model") - raise Exception( - f"Cost tracking failed for model={model}.\nDebug info - {cost_tracking_failure_debug_info}\nAdd custom pricing - https://docs.litellm.ai/docs/proxy/custom_pricing" - ) - except Exception as e: - error_msg = f"Error in tracking cost callback - {str(e)}\n Traceback:{traceback.format_exc()}" - model = kwargs.get("model", "") - metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs) - litellm_metadata = kwargs.get("litellm_params", {}).get( - "litellm_metadata", {} - ) - old_metadata = kwargs.get("litellm_params", {}).get("metadata", {}) - call_type = kwargs.get("call_type", "") - error_msg += f"\n Args to _PROXY_track_cost_callback\n model: {model}\n chosen_metadata: {metadata}\n litellm_metadata: {litellm_metadata}\n old_metadata: {old_metadata}\n call_type: {call_type}\n" - asyncio.create_task( - proxy_logging_obj.failed_tracking_alert( - error_message=error_msg, - failing_model=model, - ) - ) - - spend_log_error("Error in tracking cost callback - %s", str(e), exc=e) - - @staticmethod - async def _enrich_failure_metadata_with_key_info(metadata: dict) -> dict: - """ - Enriches failure spend log metadata by looking up the key object (and team object) - from cache/DB when key fields are missing. - - This handles two scenarios: - 1. Auth errors (401): UserAPIKeyAuth is created with only api_key set, all other - fields are null. We look up the full key object to fill in alias, user_id, - team_id, etc. - 2. Post-auth failures (provider errors, rate limits): key fields are populated - but team_alias is missing because LiteLLM_VerificationTokenView SQL view - doesn't include it. We look up the team object to fill in team_alias. - """ - api_key_hash = metadata.get("user_api_key") - if not api_key_hash: - return metadata - - from litellm.proxy.proxy_server import ( - prisma_client, - proxy_logging_obj, - user_api_key_cache, - ) - - # Step 1: If key fields are missing, look up the full key object - if metadata.get("user_api_key_alias") is None: - try: - key_obj = await get_key_object( - hashed_token=api_key_hash, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - proxy_logging_obj=proxy_logging_obj, - ) - if metadata.get("user_api_key_alias") is None: - metadata["user_api_key_alias"] = key_obj.key_alias - if metadata.get("user_api_key_user_id") is None: - metadata["user_api_key_user_id"] = key_obj.user_id - if metadata.get("user_api_key_team_id") is None: - metadata["user_api_key_team_id"] = key_obj.team_id - if metadata.get("user_api_key_org_id") is None: - metadata["user_api_key_org_id"] = key_obj.org_id - except Exception: - verbose_proxy_logger.debug( - "Failed to enrich failure metadata with key info for api_key=%s", - api_key_hash, - ) - - # Step 2: If team_id is known but team_alias is missing, look up the team object - team_id = metadata.get("user_api_key_team_id") - if team_id and metadata.get("user_api_key_team_alias") is None: - try: - team_obj = await get_team_object( - team_id=team_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - proxy_logging_obj=proxy_logging_obj, - ) - if team_obj.team_alias is not None: - metadata["user_api_key_team_alias"] = team_obj.team_alias - except Exception: - verbose_proxy_logger.debug( - "Failed to enrich failure metadata with team_alias for team_id=%s", - team_id, - ) - return metadata - - @staticmethod - def _should_track_errors_in_db(): - """ - Returns True if errors should be tracked in the database - - By default, errors are tracked in the database - - If users want to disable error tracking, they can set the disable_error_logs flag in the general_settings - """ - from litellm.proxy.proxy_server import general_settings - - if general_settings.get("disable_error_logs") is True: - return False - return - - -def _should_track_cost_callback( - user_api_key: Optional[str], - user_id: Optional[str], - team_id: Optional[str], - end_user_id: Optional[str], -) -> bool: - """ - Determine if the cost callback should be tracked based on the kwargs - """ - - # don't run track cost callback if user opted into disabling spend - if ProxyUpdateSpend.disable_spend_updates() is True: - return False - - if ( - user_api_key is not None - or user_id is not None - or team_id is not None - or end_user_id is not None - ): - return True - return False - - -def _get_budget_reservation_from_metadata(metadata: dict) -> Optional[dict]: - metadata_budget_reservation = metadata.get("user_api_key_budget_reservation") - if isinstance(metadata_budget_reservation, dict): - return metadata_budget_reservation - - user_api_key_auth_obj = metadata.get("user_api_key_auth") - if user_api_key_auth_obj is None: - return None - if isinstance(user_api_key_auth_obj, dict): - budget_reservation = user_api_key_auth_obj.get("budget_reservation") - return budget_reservation if isinstance(budget_reservation, dict) else None - return getattr(user_api_key_auth_obj, "budget_reservation", None) - - -def _get_request_tags_for_cost_tracking( - sl_object: Optional[StandardLoggingPayload], - metadata: dict, -) -> Optional[List[str]]: - if sl_object is not None: - request_tags = sl_object.get("request_tags", None) - if isinstance(request_tags, list): - return request_tags - - metadata_tags = metadata.get("tags", None) - if isinstance(metadata_tags, list): - return metadata_tags - - return None - - -async def _update_database_and_spend_counters( - proxy_logging_obj: Any, - increment_spend_counters: Any, - user_api_key: Optional[str], - user_id: Optional[str], - end_user_id: Optional[str], - team_id: Optional[str], - org_id: Optional[str], - kwargs: dict, - completion_response: Optional[Union[litellm.ModelResponse, Any]], - start_time: Any, - end_time: Any, - response_cost: float, - budget_reservation: Optional[dict], - request_tags: Optional[List[str]] = None, -) -> None: - try: - await proxy_logging_obj.db_spend_update_writer.update_database( - token=user_api_key, - response_cost=response_cost, - user_id=user_id, - end_user_id=end_user_id, - team_id=team_id, - kwargs=kwargs, - completion_response=completion_response, - start_time=start_time, - end_time=end_time, - org_id=org_id, - ) - except Exception: - if budget_reservation is not None: - try: - await _release_budget_reservation(budget_reservation=budget_reservation) - except Exception: - verbose_proxy_logger.exception( - "Failed to release budget reservation after database update failed" - ) - try: - await _invalidate_budget_reservation_counters( - budget_reservation=budget_reservation - ) - except Exception: - verbose_proxy_logger.exception( - "Failed to invalidate budget reservation counters after release failed" - ) - raise - - try: - await increment_spend_counters( - token=user_api_key, - team_id=team_id, - user_id=user_id, - response_cost=response_cost, - org_id=org_id, - budget_reservation=budget_reservation, - end_user_id=end_user_id, - tags=request_tags, - ) - except Exception: - if budget_reservation is not None: - try: - await _invalidate_budget_reservation_counters( - budget_reservation=budget_reservation - ) - except Exception: - verbose_proxy_logger.exception( - "Failed to invalidate budget reservation counters after spend counter update failed" - ) - finally: - budget_reservation["finalized"] = True - raise - - -async def _release_budget_reservation(budget_reservation: Optional[dict]) -> None: - if budget_reservation is None: - return - - from litellm.proxy.spend_tracking.budget_reservation import ( - release_budget_reservation, - ) - - await release_budget_reservation( - budget_reservation=budget_reservation, - ) - - -async def _invalidate_budget_reservation_counters( - budget_reservation: Optional[dict], -) -> None: - if budget_reservation is None: - return - - from litellm.proxy.spend_tracking.budget_reservation import ( - invalidate_budget_reservation_counters, - ) - - await invalidate_budget_reservation_counters( - budget_reservation=budget_reservation, - ) +import asyncio +import traceback +from datetime import datetime +from typing import Any, List, Optional, Union, cast + +import litellm +from litellm._logging import verbose_proxy_logger +from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.core_helpers import ( + _get_parent_otel_span_from_kwargs, + get_litellm_metadata_from_kwargs, +) +from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.auth.auth_checks import ( + get_key_object, + get_team_object, + log_db_metrics, +) +from litellm.proxy.auth.route_checks import RouteChecks +from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup +from litellm.proxy.spend_tracking.spend_log_error_logger import ( + should_suppress_spend_log_tracebacks, + spend_log_error, +) +from litellm.proxy.spend_tracking.spend_tracking_utils import ( + _sanitize_error_information_for_spend_logs, +) +from litellm.proxy.utils import ProxyUpdateSpend +from litellm.types.utils import ( + StandardLoggingPayload, + StandardLoggingPayloadErrorInformation, +) +from litellm.utils import get_end_user_id_for_cost_tracking + + +class _ProxyDBLogger(CustomLogger): + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + await self._PROXY_track_cost_callback( + kwargs, response_obj, start_time, end_time + ) + + async def async_post_call_failure_hook( # noqa: PLR0915 + self, + request_data: dict, + original_exception: Exception, + user_api_key_dict: UserAPIKeyAuth, + traceback_str: Optional[str] = None, + ): + try: + await _release_budget_reservation( + budget_reservation=user_api_key_dict.budget_reservation + ) + except Exception: + verbose_proxy_logger.exception( + "Failed to release budget reservation during failure handling" + ) + try: + await _invalidate_budget_reservation_counters( + budget_reservation=user_api_key_dict.budget_reservation + ) + if user_api_key_dict.budget_reservation is not None: + user_api_key_dict.budget_reservation["finalized"] = True + except Exception: + verbose_proxy_logger.exception( + "Failed to invalidate budget reservation counters after failure release failed" + ) + + request_route = user_api_key_dict.request_route + if _ProxyDBLogger._should_track_errors_in_db() is False: + return + elif request_route is not None and not ( + RouteChecks.is_llm_api_route(route=request_route) + or RouteChecks.is_info_route(route=request_route) + ): + return + + from litellm.proxy.proxy_server import proxy_logging_obj + + _metadata = dict( + LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key( + user_api_key_dict=user_api_key_dict + ) + ) + _metadata["user_api_key"] = user_api_key_dict.api_key + _metadata["status"] = "failure" + _error_information = StandardLoggingPayloadSetup.get_error_information( + original_exception=original_exception, + traceback_str=traceback_str, + ) + if should_suppress_spend_log_tracebacks(): + # Drop the traceback key entirely so the per-row Metadata pane in + # the UI (which renders the JSON blob verbatim) doesn't show a + # noisy ``"traceback": ""`` line. Downstream consumers all use + # ``.get("traceback")`` / truthy checks, and the TypedDict marks + # the field as optional, so omitting is type-safe. + _error_information.pop("traceback", None) + # Strip echoed request input + apply DB-size cap before storing in + # the spend-log metadata column (LIT-2992). Result is never None + # here because the input above is constructed non-None. + _error_information = cast( + StandardLoggingPayloadErrorInformation, + _sanitize_error_information_for_spend_logs(_error_information), + ) + _metadata["error_information"] = _error_information + + _metadata = await _ProxyDBLogger._enrich_failure_metadata_with_key_info( + metadata=_metadata, + ) + + existing_metadata: dict = request_data.get("metadata", None) or {} + existing_metadata.update(_metadata) + + if "litellm_params" not in request_data: + request_data["litellm_params"] = {} + + existing_litellm_params = request_data.get("litellm_params", {}) + existing_litellm_metadata = existing_litellm_params.get("metadata", {}) or {} + + # Preserve tags from existing metadata + if existing_litellm_metadata.get("tags"): + existing_metadata["tags"] = existing_litellm_metadata.get("tags") + + request_data["litellm_params"]["proxy_server_request"] = ( + request_data.get("proxy_server_request") + or existing_litellm_params.get("proxy_server_request") + or {} + ) + request_data["litellm_params"]["metadata"] = existing_metadata + + # Preserve model name and custom_llm_provider + if "model" not in request_data: + request_data["model"] = existing_litellm_params.get( + "model" + ) or request_data.get("model", "") + if "custom_llm_provider" not in request_data: + request_data["custom_llm_provider"] = existing_litellm_params.get( + "custom_llm_provider" + ) or request_data.get("custom_llm_provider", "") + + # Propagate standard_logging_object and litellm_trace_id from the + # Logging instance so that _get_session_id_for_spend_log uses the same + # trace_id that Langfuse received (via async_failure_handler). + # Without this, the DB session_id would be a random UUID that doesn't + # match the Langfuse trace_id, making failed requests unsearchable. + _litellm_logging_obj = request_data.get("litellm_logging_obj") + if _litellm_logging_obj is not None: + if not request_data.get("standard_logging_object"): + request_data["standard_logging_object"] = getattr( + _litellm_logging_obj, "model_call_details", {} + ).get("standard_logging_object") + if request_data.get("litellm_trace_id") is None: + request_data["litellm_trace_id"] = getattr( + _litellm_logging_obj, "litellm_trace_id", None + ) + + # Use the actual request start time from the logging object so that + # failed requests record the real duration instead of 0. + actual_start_time = datetime.now() + if _litellm_logging_obj is not None: + obj_start = getattr(_litellm_logging_obj, "start_time", None) + if obj_start is not None: + actual_start_time = obj_start + + # A stream that broke mid-flight still billed the provider for the + # chunks already delivered. ``post_call_failure_hook`` lifts that + # recovered cost onto request_data (the usage rides along in + # ``combined_usage_object`` for the token columns), so attribute the + # real partial spend to this failure row instead of zero. + recovered_response_cost = 0.0 + if isinstance(request_data.get("combined_usage_object"), litellm.Usage): + recovered_response_cost = max( + float(request_data.get("response_cost") or 0.0), 0.0 + ) + + await proxy_logging_obj.db_spend_update_writer.update_database( + token=user_api_key_dict.api_key, + response_cost=recovered_response_cost, + user_id=user_api_key_dict.user_id, + end_user_id=user_api_key_dict.end_user_id, + team_id=user_api_key_dict.team_id, + kwargs=request_data, + completion_response=original_exception, + start_time=actual_start_time, + end_time=datetime.now(), + org_id=user_api_key_dict.org_id, + ) + + @log_db_metrics + async def _PROXY_track_cost_callback( + self, + kwargs, # kwargs to completion + completion_response: Optional[ + Union[litellm.ModelResponse, Any] + ], # response from completion + start_time=None, + end_time=None, # start/end time for completion + ): + from litellm.proxy.proxy_server import ( + increment_spend_counters, + proxy_logging_obj, + update_cache, + ) + + verbose_proxy_logger.debug("INSIDE _PROXY_track_cost_callback") + try: + verbose_proxy_logger.debug( + f"kwargs stream: {kwargs.get('stream', None)} + complete streaming response: {kwargs.get('complete_streaming_response', None)}" + ) + parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs=kwargs) + litellm_params = kwargs.get("litellm_params", {}) or {} + end_user_id = get_end_user_id_for_cost_tracking(litellm_params) + metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs) + budget_reservation = _get_budget_reservation_from_metadata( + metadata=metadata + ) + user_id = cast(Optional[str], metadata.get("user_api_key_user_id", None)) + team_id = cast(Optional[str], metadata.get("user_api_key_team_id", None)) + org_id = cast(Optional[str], metadata.get("user_api_key_org_id", None)) + key_alias = cast(Optional[str], metadata.get("user_api_key_alias", None)) + end_user_max_budget = metadata.get("user_api_end_user_max_budget", None) + sl_object: Optional[StandardLoggingPayload] = kwargs.get( + "standard_logging_object", None + ) + response_cost = ( + sl_object.get("response_cost", None) + if sl_object is not None + else kwargs.get("response_cost", None) + ) + tags = _get_request_tags_for_cost_tracking( + sl_object=sl_object, + metadata=metadata, + ) + + if response_cost is not None: + user_api_key = metadata.get("user_api_key", None) + if kwargs.get("cache_hit", False) is True: + response_cost = 0.0 + verbose_proxy_logger.debug( + f"Cache Hit: response_cost {response_cost}, for user_id {user_id}" + ) + + verbose_proxy_logger.debug( + f"user_api_key {user_api_key}, user_id {user_id}, team_id {team_id}, end_user_id {end_user_id}" + ) + if _should_track_cost_callback( + user_api_key=user_api_key, + user_id=user_id, + team_id=team_id, + end_user_id=end_user_id, + ): + ## UPDATE DATABASE + await _update_database_and_spend_counters( + proxy_logging_obj=proxy_logging_obj, + increment_spend_counters=increment_spend_counters, + user_api_key=user_api_key, + user_id=user_id, + end_user_id=end_user_id, + team_id=team_id, + org_id=org_id, + kwargs=kwargs, + completion_response=completion_response, + start_time=start_time, + end_time=end_time, + response_cost=response_cost, + budget_reservation=budget_reservation, + request_tags=tags, + ) + + # update cache (fire-and-forget for backward compat: + # cached object fields, soft budget alerts, etc.) + asyncio.create_task( + update_cache( + token=user_api_key, + user_id=user_id, + end_user_id=end_user_id, + response_cost=response_cost, + team_id=team_id, + parent_otel_span=parent_otel_span, + tags=tags, + ) + ) + + await proxy_logging_obj.slack_alerting_instance.customer_spend_alert( + token=user_api_key, + key_alias=key_alias, + end_user_id=end_user_id, + response_cost=response_cost, + max_budget=end_user_max_budget, + ) + elif budget_reservation is not None: + await _release_budget_reservation( + budget_reservation=budget_reservation + ) + else: + await _release_budget_reservation(budget_reservation=budget_reservation) + # Non-model call types (health checks, afile_delete) have no model or standard_logging_object. + # Use .get() for "stream" to avoid KeyError on health checks. + if sl_object is None and not kwargs.get("model"): + verbose_proxy_logger.warning( + "Cost tracking - skipping, no standard_logging_object and no model for call_type=%s", + kwargs.get("call_type", "unknown"), + ) + return + if kwargs.get("stream") is not True or ( + kwargs.get("stream") is True + and "complete_streaming_response" in kwargs + ): + if sl_object is not None: + cost_tracking_failure_debug_info: Union[dict, str] = ( + sl_object["response_cost_failure_debug_info"] # type: ignore + or "response_cost_failure_debug_info is None in standard_logging_object" + ) + else: + cost_tracking_failure_debug_info = ( + "standard_logging_object not found" + ) + model = kwargs.get("model") + raise Exception( + f"Cost tracking failed for model={model}.\nDebug info - {cost_tracking_failure_debug_info}\nAdd custom pricing - https://docs.litellm.ai/docs/proxy/custom_pricing" + ) + except Exception as e: + error_msg = f"Error in tracking cost callback - {str(e)}\n Traceback:{traceback.format_exc()}" + model = kwargs.get("model", "") + metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs) + litellm_metadata = kwargs.get("litellm_params", {}).get( + "litellm_metadata", {} + ) + old_metadata = kwargs.get("litellm_params", {}).get("metadata", {}) + call_type = kwargs.get("call_type", "") + error_msg += f"\n Args to _PROXY_track_cost_callback\n model: {model}\n chosen_metadata: {metadata}\n litellm_metadata: {litellm_metadata}\n old_metadata: {old_metadata}\n call_type: {call_type}\n" + asyncio.create_task( + proxy_logging_obj.failed_tracking_alert( + error_message=error_msg, + failing_model=model, + ) + ) + + spend_log_error("Error in tracking cost callback - %s", str(e), exc=e) + + @staticmethod + async def _enrich_failure_metadata_with_key_info(metadata: dict) -> dict: + """ + Enriches failure spend log metadata by looking up the key object (and team object) + from cache/DB when key fields are missing. + + This handles two scenarios: + 1. Auth errors (401): UserAPIKeyAuth is created with only api_key set, all other + fields are null. We look up the full key object to fill in alias, user_id, + team_id, etc. + 2. Post-auth failures (provider errors, rate limits): key fields are populated + but team_alias is missing because LiteLLM_VerificationTokenView SQL view + doesn't include it. We look up the team object to fill in team_alias. + """ + api_key_hash = metadata.get("user_api_key") + if not api_key_hash: + return metadata + + from litellm.proxy.proxy_server import ( + prisma_client, + proxy_logging_obj, + user_api_key_cache, + ) + + # Step 1: If key fields are missing, look up the full key object + if metadata.get("user_api_key_alias") is None: + try: + key_obj = await get_key_object( + hashed_token=api_key_hash, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + if metadata.get("user_api_key_alias") is None: + metadata["user_api_key_alias"] = key_obj.key_alias + if metadata.get("user_api_key_user_id") is None: + metadata["user_api_key_user_id"] = key_obj.user_id + if metadata.get("user_api_key_team_id") is None: + metadata["user_api_key_team_id"] = key_obj.team_id + if metadata.get("user_api_key_org_id") is None: + metadata["user_api_key_org_id"] = key_obj.org_id + except Exception: + verbose_proxy_logger.debug( + "Failed to enrich failure metadata with key info for api_key=%s", + api_key_hash, + ) + + # Step 2: If team_id is known but team_alias is missing, look up the team object + team_id = metadata.get("user_api_key_team_id") + if team_id and metadata.get("user_api_key_team_alias") is None: + try: + team_obj = await get_team_object( + team_id=team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + if team_obj.team_alias is not None: + metadata["user_api_key_team_alias"] = team_obj.team_alias + except Exception: + verbose_proxy_logger.debug( + "Failed to enrich failure metadata with team_alias for team_id=%s", + team_id, + ) + return metadata + + @staticmethod + def _should_track_errors_in_db(): + """ + Returns True if errors should be tracked in the database + + By default, errors are tracked in the database + + If users want to disable error tracking, they can set the disable_error_logs flag in the general_settings + """ + from litellm.proxy.proxy_server import general_settings + + if general_settings.get("disable_error_logs") is True: + return False + return + + +def _should_track_cost_callback( + user_api_key: Optional[str], + user_id: Optional[str], + team_id: Optional[str], + end_user_id: Optional[str], +) -> bool: + """ + Determine if the cost callback should be tracked based on the kwargs + """ + + # don't run track cost callback if user opted into disabling spend + if ProxyUpdateSpend.disable_spend_updates() is True: + return False + + if ( + user_api_key is not None + or user_id is not None + or team_id is not None + or end_user_id is not None + ): + return True + return False + + +def _get_budget_reservation_from_metadata(metadata: dict) -> Optional[dict]: + metadata_budget_reservation = metadata.get("user_api_key_budget_reservation") + if isinstance(metadata_budget_reservation, dict): + return metadata_budget_reservation + + user_api_key_auth_obj = metadata.get("user_api_key_auth") + if user_api_key_auth_obj is None: + return None + if isinstance(user_api_key_auth_obj, dict): + budget_reservation = user_api_key_auth_obj.get("budget_reservation") + return budget_reservation if isinstance(budget_reservation, dict) else None + return getattr(user_api_key_auth_obj, "budget_reservation", None) + + +def _get_request_tags_for_cost_tracking( + sl_object: Optional[StandardLoggingPayload], + metadata: dict, +) -> Optional[List[str]]: + if sl_object is not None: + request_tags = sl_object.get("request_tags", None) + if isinstance(request_tags, list): + return request_tags + + metadata_tags = metadata.get("tags", None) + if isinstance(metadata_tags, list): + return metadata_tags + + return None + + +async def _update_database_and_spend_counters( + proxy_logging_obj: Any, + increment_spend_counters: Any, + user_api_key: Optional[str], + user_id: Optional[str], + end_user_id: Optional[str], + team_id: Optional[str], + org_id: Optional[str], + kwargs: dict, + completion_response: Optional[Union[litellm.ModelResponse, Any]], + start_time: Any, + end_time: Any, + response_cost: float, + budget_reservation: Optional[dict], + request_tags: Optional[List[str]] = None, +) -> None: + try: + await proxy_logging_obj.db_spend_update_writer.update_database( + token=user_api_key, + response_cost=response_cost, + user_id=user_id, + end_user_id=end_user_id, + team_id=team_id, + kwargs=kwargs, + completion_response=completion_response, + start_time=start_time, + end_time=end_time, + org_id=org_id, + ) + except Exception: + if budget_reservation is not None: + try: + await _release_budget_reservation(budget_reservation=budget_reservation) + except Exception: + verbose_proxy_logger.exception( + "Failed to release budget reservation after database update failed" + ) + try: + await _invalidate_budget_reservation_counters( + budget_reservation=budget_reservation + ) + except Exception: + verbose_proxy_logger.exception( + "Failed to invalidate budget reservation counters after release failed" + ) + raise + + try: + await increment_spend_counters( + token=user_api_key, + team_id=team_id, + user_id=user_id, + response_cost=response_cost, + org_id=org_id, + budget_reservation=budget_reservation, + end_user_id=end_user_id, + tags=request_tags, + ) + except Exception: + if budget_reservation is not None: + try: + await _invalidate_budget_reservation_counters( + budget_reservation=budget_reservation + ) + except Exception: + verbose_proxy_logger.exception( + "Failed to invalidate budget reservation counters after spend counter update failed" + ) + finally: + budget_reservation["finalized"] = True + raise + + +async def _release_budget_reservation(budget_reservation: Optional[dict]) -> None: + if budget_reservation is None: + return + + from litellm.proxy.spend_tracking.budget_reservation import ( + release_budget_reservation, + ) + + await release_budget_reservation( + budget_reservation=budget_reservation, + ) + + +async def _invalidate_budget_reservation_counters( + budget_reservation: Optional[dict], +) -> None: + if budget_reservation is None: + return + + from litellm.proxy.spend_tracking.budget_reservation import ( + invalidate_budget_reservation_counters, + ) + + await invalidate_budget_reservation_counters( + budget_reservation=budget_reservation, + ) diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index e2881faca0d..26d8d748d2e 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -265,6 +265,13 @@ def get_logging_payload( # noqa: PLR0915 elif isinstance(_usage, dict): usage = _usage + # A request that failed mid-stream has no usable response_obj usage, but the + # streaming handler may have recovered the usage from the chunks already + # delivered. Honor that override so the partial usage lands in spend tracking. + _combined_usage = kwargs.get("combined_usage_object") + if not usage and isinstance(_combined_usage, litellm.Usage): + usage = _combined_usage.model_dump() + id = get_spend_logs_id(call_type or "acompletion", response_obj_dict, kwargs) standard_logging_payload = cast( Optional[StandardLoggingPayload], kwargs.get("standard_logging_object", None) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 9cce9eb3812..e5af9b482cd 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -1826,6 +1826,20 @@ class ProxyLogging: original_exception=original_exception, ) + _logging_obj = request_data.get("litellm_logging_obj") + if _logging_obj is not None: + _model_call_details = getattr(_logging_obj, "model_call_details", {}) + + # A stream that broke mid-flight still billed the provider for the + # chunks already delivered; the streaming handler stashes that + # recovered usage and cost here. Lift them onto request_data so the + # failure-path spend callbacks (which run after the logging object + # is popped) record the real partial spend instead of zero. + _recovered_usage = _model_call_details.get("combined_usage_object") + if _recovered_usage is not None: + request_data["combined_usage_object"] = _recovered_usage + request_data["response_cost"] = _model_call_details.get("response_cost") + # Remove before callbacks iterate — not serialisable request_data.pop("litellm_logging_obj", None) 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 c4efd63ec1e..1fc3897a4de 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -2891,3 +2891,46 @@ def test_success_handler_unified_helper_runs_for_typed_results(): ) mock_calc.assert_called_once() assert logging_obj.model_call_details["response_cost"] == expected_cost + + +def test_failure_handler_records_recovered_partial_spend(logging_obj): + """A stream interrupted mid-flight still billed the provider for the chunks + already delivered. When the router stashes that recovered usage as + ``combined_usage_object`` and pre-computes ``response_cost``, the failure + handler must preserve them so the failure row carries the real partial + spend instead of zero. + """ + from litellm.types.utils import Usage + + logging_obj.model_call_details["combined_usage_object"] = Usage( + prompt_tokens=17, completion_tokens=9, total_tokens=26 + ) + logging_obj.model_call_details["response_cost"] = 0.00012 + + logging_obj._failure_handler_helper_fn( + exception=Exception("Connection lost"), + traceback_exception="Traceback ...", + ) + + payload = logging_obj.model_call_details["standard_logging_object"] + assert payload["status"] == "failure" + assert payload["response_cost"] == 0.00012 + assert payload["prompt_tokens"] == 17 + assert payload["completion_tokens"] == 9 + assert payload["total_tokens"] == 26 + + +def test_failure_handler_zeroes_spend_without_recovered_usage(logging_obj): + """A failure with no recovered partial usage keeps the existing behavior of + recording zero spend, so the partial-spend preservation does not leak into + ordinary failures. + """ + logging_obj._failure_handler_helper_fn( + exception=Exception("boom"), + traceback_exception="Traceback ...", + ) + + payload = logging_obj.model_call_details["standard_logging_object"] + assert payload["status"] == "failure" + assert payload["response_cost"] == 0 + assert payload["total_tokens"] == 0 diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py index 63e2cb7f35c..09bb8532ec2 100644 --- a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py +++ b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py @@ -2118,3 +2118,79 @@ def test_gemini_legacy_vertex_tool_calls_finish_reason_with_stop_enum(): f"Expected 'tool_calls' but got {final.choices[0].finish_reason!r}. " "STOP enum was not normalised through map_finish_reason()." ) + + +def test_record_partial_usage_for_failure_stashes_usage_and_cost(): + """A stream that breaks mid-flight must surface the usage assembled from the + chunks already delivered, plus its cost, on the logging object so the + failure handler records the real partial spend instead of zero. + """ + logging_obj = Logging( + model="gpt-4o-mini", + messages=[{"role": "user", "content": "Hey"}], + stream=True, + call_type="completion", + start_time=time.time(), + litellm_call_id="partial-usage-1", + function_id="1245", + ) + logging_obj.model_call_details["custom_llm_provider"] = "openai" + + wrapper = CustomStreamWrapper( + completion_stream=None, + model="gpt-4o-mini", + logging_obj=logging_obj, + custom_llm_provider="openai", + ) + wrapper.chunks = [ + ModelResponseStream( + id="chatcmpl-partial-1", + created=1742056047, + model="gpt-4o-mini", + object="chat.completion.chunk", + choices=[ + StreamingChoices( + finish_reason=None, + index=0, + delta=Delta( + content="The Roman Empire began when", role="assistant" + ), + ) + ], + usage=Usage(prompt_tokens=30, completion_tokens=1, total_tokens=31), + ) + ] + + wrapper._record_partial_usage_for_failure() + + stashed = logging_obj.model_call_details["combined_usage_object"] + assert stashed.prompt_tokens == 30 + assert stashed.completion_tokens == 1 + assert stashed.total_tokens == 31 + assert isinstance(logging_obj.model_call_details["response_cost"], float) + + +def test_record_partial_usage_for_failure_noop_without_chunks(): + """With no chunks delivered there is nothing billed to recover, so the + failure stash must stay absent and not force a zero-usage row. + """ + logging_obj = Logging( + model="gpt-4o-mini", + messages=[{"role": "user", "content": "Hey"}], + stream=True, + call_type="completion", + start_time=time.time(), + litellm_call_id="partial-usage-2", + function_id="1245", + ) + wrapper = CustomStreamWrapper( + completion_stream=None, + model="gpt-4o-mini", + logging_obj=logging_obj, + custom_llm_provider="openai", + ) + wrapper.chunks = [] + + wrapper._record_partial_usage_for_failure() + + assert "combined_usage_object" not in logging_obj.model_call_details diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index 771e10a54a0..0cbf308076c 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -1067,3 +1067,40 @@ async def test_failure_hook_drops_error_information_traceback_when_env_set( assert "traceback" not in error_information assert error_information["error_class"] == "RuntimeError" assert error_information["error_message"] == "boom-with-traceback" + + +@pytest.mark.asyncio +async def test_async_post_call_failure_hook_records_recovered_partial_spend(): + """A stream that broke mid-flight still billed the provider. The failure + hook lifts the recovered cost onto request_data as ``response_cost``; this + hook must pass it through to update_database so the failure row records the + real partial spend instead of the hardcoded zero. + """ + from litellm.types.utils import Usage + + logger = _ProxyDBLogger() + user_api_key_dict = UserAPIKeyAuth(api_key="test_api_key", user_id="u", team_id="t") + + request_data = { + "model": "anthropic/claude-haiku-4-5", + "messages": [{"role": "user", "content": "Hello"}], + "metadata": {}, + "proxy_server_request": {"request_id": "rid"}, + "response_cost": 3.5e-05, + "combined_usage_object": Usage( + prompt_tokens=30, completion_tokens=1, total_tokens=31 + ), + } + + with patch( + "litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database", + new_callable=AsyncMock, + ) as mock_update_database: + await logger.async_post_call_failure_hook( + request_data=request_data, + original_exception=Exception("MidStreamFallbackError: read timeout"), + user_api_key_dict=user_api_key_dict, + ) + + mock_update_database.assert_called_once() + assert mock_update_database.call_args[1]["response_cost"] == 3.5e-05 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 5ca058fc8d9..5272b105eb5 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 @@ -2009,3 +2009,50 @@ def test_sanitize_error_information_redacts_pydantic_assignment_form( assert sanitized is not None assert "leaked-via-pydantic-msg" not in sanitized["error_message"] assert REDACTED_BY_LITELM_STRING in sanitized["error_message"] + + +def test_get_logging_payload_uses_recovered_combined_usage_on_failure(): + """A request that fails mid-stream has no usable response_obj usage, but the + streaming handler recovers the usage from the chunks already delivered and + the failure hook surfaces it as ``combined_usage_object``. The spend-log + payload must record those token counts instead of zero. + """ + from litellm.types.utils import Usage + + kwargs = { + "model": "anthropic/claude-haiku-4-5", + "call_type": "acompletion", + "litellm_params": {"metadata": {"user_api_key": "sk-test"}}, + "combined_usage_object": Usage( + prompt_tokens=30, completion_tokens=1, total_tokens=31 + ), + } + response_obj = Exception("MidStreamFallbackError: read timeout") + now = datetime.datetime.now(timezone.utc) + + payload = get_logging_payload( + kwargs=kwargs, response_obj=response_obj, start_time=now, end_time=now + ) + + assert payload["prompt_tokens"] == 30 + assert payload["completion_tokens"] == 1 + assert payload["total_tokens"] == 31 + + +def test_get_logging_payload_failure_without_recovered_usage_is_zero(): + """A failure with no recovered usage keeps zero token counts, so the + combined-usage override never invents tokens for ordinary failures. + """ + kwargs = { + "model": "anthropic/claude-haiku-4-5", + "call_type": "acompletion", + "litellm_params": {"metadata": {"user_api_key": "sk-test"}}, + } + response_obj = Exception("BadRequestError") + now = datetime.datetime.now(timezone.utc) + + payload = get_logging_payload( + kwargs=kwargs, response_obj=response_obj, start_time=now, end_time=now + ) + + assert payload["total_tokens"] == 0 diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index 2605eadba7a..5a21f2c0f2c 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -264,3 +264,52 @@ def test_enrich_http_exception_callback_without_guardrail_name_noop(): exc = HTTPException(status_code=400, detail={"error": "x"}) _enrich_http_exception_with_guardrail_context(exc, StubCallback()) assert exc.detail == {"error": "x"} + + +class TestPostCallFailureHookLiftsRecoveredPartialSpend: + """A stream that broke mid-flight still billed the provider for the chunks + already delivered. The streaming handler stashes that recovered usage and + cost on the logging object; post_call_failure_hook must lift them onto + request_data before the logging object is popped, so the failure-path spend + callbacks (which run after the pop) record the real partial spend. + """ + + async def _run(self, request_data): + from unittest.mock import AsyncMock, patch + + from litellm.proxy._types import UserAPIKeyAuth + + proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + proxy_logging_obj.alert_types = [] + with patch.object(proxy_logging_obj, "update_request_status", new=AsyncMock()): + await proxy_logging_obj.post_call_failure_hook( + request_data=request_data, + original_exception=Exception("boom"), + user_api_key_dict=UserAPIKeyAuth(), + ) + + @pytest.mark.asyncio + async def test_lifts_recovered_usage_and_cost(self): + from litellm.types.utils import Usage + + recovered_usage = Usage(prompt_tokens=30, completion_tokens=1, total_tokens=31) + logging_obj = MagicMock() + logging_obj.model_call_details = { + "combined_usage_object": recovered_usage, + "response_cost": 3.5e-05, + } + request_data = {"litellm_logging_obj": logging_obj, "metadata": {}} + await self._run(request_data) + + assert request_data["combined_usage_object"] is recovered_usage + assert request_data["response_cost"] == 3.5e-05 + assert "litellm_logging_obj" not in request_data + + @pytest.mark.asyncio + async def test_no_recovered_usage_is_noop(self): + logging_obj = MagicMock() + logging_obj.model_call_details = {} + request_data = {"litellm_logging_obj": logging_obj, "metadata": {}} + await self._run(request_data) + assert "combined_usage_object" not in request_data + assert "response_cost" not in request_data diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 48facace528..cbaa645b964 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -2984,6 +2984,151 @@ def test_combine_fallback_usage(): assert chunk.usage.total_tokens == 15 +@pytest.mark.asyncio +async def test_acompletion_streaming_iterator_does_not_log_success_on_terminal_failure(): + """A mid-stream failure with no successful fallback raises and is logged as + a failure, so the router must never dispatch it as a success. Partial-spend + recovery for the failure row happens in the streaming handler, not here, so + this guards only against reintroducing a success log for a failed stream. + """ + from litellm.exceptions import MidStreamFallbackError + from litellm.types.utils import Delta, StreamingChoices, Usage + + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4", "api_key": "fake-key-1"}, + }, + ], + set_verbose=True, + ) + + error = MidStreamFallbackError( + message="Connection lost", + model="gpt-4", + llm_provider="openai", + generated_content="The Roman Empire began when", + ) + + def _make_interrupted_model_response(): + partial_chunk = litellm.ModelResponseStream( + id="chatcmpl-partial-1", + created=1742056047, + model="gpt-4", + object="chat.completion.chunk", + choices=[ + StreamingChoices( + finish_reason=None, + index=0, + delta=Delta( + content="The Roman Empire began when", role="assistant" + ), + ) + ], + usage=Usage(prompt_tokens=17, completion_tokens=9, total_tokens=26), + ) + + class _RaisingStream: + def __init__(self): + self.index = 0 + self.chunks = [partial_chunk] + + def __aiter__(self): + return self + + async def __anext__(self): + if self.index == 0: + self.index += 1 + return partial_chunk + raise error + + stream = _RaisingStream() + logging_obj = MagicMock() + logging_obj.dispatch_success_handlers = AsyncMock() + logging_obj.model_call_details = {} + setattr(stream, "model", "gpt-4") + setattr(stream, "custom_llm_provider", "openai") + setattr(stream, "logging_obj", logging_obj) + return stream, logging_obj + + messages = [{"role": "user", "content": "Hello"}] + initial_kwargs = {"model": "gpt-4", "stream": True} + + # Terminal path: no successful fallback -> the error propagates and the + # router never dispatches a success for the failed stream. + model_response, logging_obj = _make_interrupted_model_response() + with patch.object( + router, + "async_function_with_fallbacks_common_utils", + new=AsyncMock(side_effect=error), + ): + result = await router._acompletion_streaming_iterator( + model_response=model_response, + messages=messages, + initial_kwargs=dict(initial_kwargs), + ) + collected = [] + with pytest.raises(MidStreamFallbackError): + async for chunk in result: + collected.append(chunk) + + assert len(collected) == 1 + logging_obj.dispatch_success_handlers.assert_not_called() + + # Fallback success: the fallback stream owns success accounting via + # _combine_fallback_usage, so this iterator must not dispatch its own. + model_response, logging_obj = _make_interrupted_model_response() + + class _FallbackStream: + def __init__(self, items): + self.items = items + self.index = 0 + + def __aiter__(self): + return self + + async def __anext__(self): + if self.index >= len(self.items): + raise StopAsyncIteration + item = self.items[self.index] + self.index += 1 + return item + + fallback_stream = _FallbackStream( + [ + litellm.ModelResponseStream( + id="chatcmpl-fallback-1", + model="gpt-3.5-turbo", + object="chat.completion.chunk", + choices=[ + StreamingChoices( + finish_reason=None, + index=0, + delta=Delta(content=" continued", role="assistant"), + ) + ], + ) + ] + ) + with patch.object( + router, + "async_function_with_fallbacks_common_utils", + new=AsyncMock(return_value=fallback_stream), + ): + result = await router._acompletion_streaming_iterator( + model_response=model_response, + messages=messages, + initial_kwargs=dict(initial_kwargs), + ) + collected = [] + async for chunk in result: + collected.append(chunk) + + assert len(collected) == 2 + logging_obj.dispatch_success_handlers.assert_not_called() + + @pytest.mark.asyncio async def test_team_scoped_model_fallback(): """