fix mypy error

This commit is contained in:
Harshit Jain
2026-02-19 14:04:00 +05:30
parent 31752c7b78
commit bdf01fa283
4 changed files with 8797 additions and 8784 deletions
@@ -1,309 +1,311 @@
""" """
PagerDuty Alerting Integration PagerDuty Alerting Integration
Handles two types of alerts: Handles two types of alerts:
- High LLM API Failure Rate. Configure X fails in Y seconds to trigger an alert. - High LLM API Failure Rate. Configure X fails in Y seconds to trigger an alert.
- High Number of Hanging LLM Requests. Configure X hangs in Y seconds to trigger an alert. - High Number of Hanging LLM Requests. Configure X hangs in Y seconds to trigger an alert.
Note: This is a Free feature on the regular litellm docker image. Note: This is a Free feature on the regular litellm docker image.
However, this is under the enterprise license However, this is under the enterprise license
""" """
import asyncio import asyncio
import os import os
from datetime import datetime, timedelta, timezone from datetime import datetime, timedelta, timezone
from typing import List, Literal, Optional, Union from typing import List, Optional, Union
from litellm._logging import verbose_logger from litellm._logging import verbose_logger
from litellm.caching import DualCache from litellm.caching import DualCache
from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
from litellm.llms.custom_httpx.http_handler import ( from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler, AsyncHTTPHandler,
get_async_httpx_client, get_async_httpx_client,
httpxSpecialProvider, httpxSpecialProvider,
) )
from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.integrations.pagerduty import ( from litellm.types.integrations.pagerduty import (
AlertingConfig, AlertingConfig,
PagerDutyInternalEvent, PagerDutyInternalEvent,
PagerDutyPayload, PagerDutyPayload,
PagerDutyRequestBody, PagerDutyRequestBody,
) )
from litellm.types.utils import ( from litellm.types.utils import (
CallTypesLiteral, CallTypesLiteral,
StandardLoggingPayload, StandardLoggingPayload,
StandardLoggingPayloadErrorInformation, StandardLoggingPayloadErrorInformation,
) )
PAGERDUTY_DEFAULT_FAILURE_THRESHOLD = 60 PAGERDUTY_DEFAULT_FAILURE_THRESHOLD = 60
PAGERDUTY_DEFAULT_FAILURE_THRESHOLD_WINDOW_SECONDS = 60 PAGERDUTY_DEFAULT_FAILURE_THRESHOLD_WINDOW_SECONDS = 60
PAGERDUTY_DEFAULT_HANGING_THRESHOLD_SECONDS = 60 PAGERDUTY_DEFAULT_HANGING_THRESHOLD_SECONDS = 60
PAGERDUTY_DEFAULT_HANGING_THRESHOLD_WINDOW_SECONDS = 600 PAGERDUTY_DEFAULT_HANGING_THRESHOLD_WINDOW_SECONDS = 600
class PagerDutyAlerting(SlackAlerting): class PagerDutyAlerting(SlackAlerting):
""" """
Tracks failed requests and hanging requests separately. Tracks failed requests and hanging requests separately.
If threshold is crossed for either type, triggers a PagerDuty alert. If threshold is crossed for either type, triggers a PagerDuty alert.
""" """
def __init__( def __init__(
self, alerting_args: Optional[Union[AlertingConfig, dict]] = None, **kwargs self, alerting_args: Optional[Union[AlertingConfig, dict]] = None, **kwargs
): ):
super().__init__() super().__init__()
_api_key = os.getenv("PAGERDUTY_API_KEY") _api_key = os.getenv("PAGERDUTY_API_KEY")
if not _api_key: if not _api_key:
raise ValueError("PAGERDUTY_API_KEY is not set") raise ValueError("PAGERDUTY_API_KEY is not set")
self.api_key: str = _api_key self.api_key: str = _api_key
alerting_args = alerting_args or {} alerting_args = alerting_args or {}
self.pagerduty_alerting_args: AlertingConfig = AlertingConfig( self.pagerduty_alerting_args: AlertingConfig = AlertingConfig(
failure_threshold=alerting_args.get( failure_threshold=alerting_args.get(
"failure_threshold", PAGERDUTY_DEFAULT_FAILURE_THRESHOLD "failure_threshold", PAGERDUTY_DEFAULT_FAILURE_THRESHOLD
), ),
failure_threshold_window_seconds=alerting_args.get( failure_threshold_window_seconds=alerting_args.get(
"failure_threshold_window_seconds", "failure_threshold_window_seconds",
PAGERDUTY_DEFAULT_FAILURE_THRESHOLD_WINDOW_SECONDS, PAGERDUTY_DEFAULT_FAILURE_THRESHOLD_WINDOW_SECONDS,
), ),
hanging_threshold_seconds=alerting_args.get( hanging_threshold_seconds=alerting_args.get(
"hanging_threshold_seconds", PAGERDUTY_DEFAULT_HANGING_THRESHOLD_SECONDS "hanging_threshold_seconds", PAGERDUTY_DEFAULT_HANGING_THRESHOLD_SECONDS
), ),
hanging_threshold_window_seconds=alerting_args.get( hanging_threshold_window_seconds=alerting_args.get(
"hanging_threshold_window_seconds", "hanging_threshold_window_seconds",
PAGERDUTY_DEFAULT_HANGING_THRESHOLD_WINDOW_SECONDS, PAGERDUTY_DEFAULT_HANGING_THRESHOLD_WINDOW_SECONDS,
), ),
) )
# Separate storage for failures vs. hangs # Separate storage for failures vs. hangs
self._failure_events: List[PagerDutyInternalEvent] = [] self._failure_events: List[PagerDutyInternalEvent] = []
self._hanging_events: List[PagerDutyInternalEvent] = [] self._hanging_events: List[PagerDutyInternalEvent] = []
# ------------------ MAIN LOGIC ------------------ # # ------------------ MAIN LOGIC ------------------ #
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
""" """
Record a failure event. Only send an alert to PagerDuty if the Record a failure event. Only send an alert to PagerDuty if the
configured *failure* threshold is exceeded in the specified window. configured *failure* threshold is exceeded in the specified window.
""" """
now = datetime.now(timezone.utc) now = datetime.now(timezone.utc)
standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get( standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get(
"standard_logging_object" "standard_logging_object"
) )
if not standard_logging_payload: if not standard_logging_payload:
raise ValueError( raise ValueError(
"standard_logging_object is required for PagerDutyAlerting" "standard_logging_object is required for PagerDutyAlerting"
) )
# Extract error details # Extract error details
error_info: Optional[StandardLoggingPayloadErrorInformation] = ( error_info: Optional[StandardLoggingPayloadErrorInformation] = (
standard_logging_payload.get("error_information") or {} standard_logging_payload.get("error_information") or {}
) )
_meta = standard_logging_payload.get("metadata") or {} _meta = standard_logging_payload.get("metadata") or {}
self._failure_events.append( self._failure_events.append(
PagerDutyInternalEvent( PagerDutyInternalEvent(
failure_event_type="failed_response", failure_event_type="failed_response",
timestamp=now, timestamp=now,
error_class=error_info.get("error_class"), error_class=error_info.get("error_class"),
error_code=error_info.get("error_code"), error_code=error_info.get("error_code"),
error_llm_provider=error_info.get("llm_provider"), error_llm_provider=error_info.get("llm_provider"),
user_api_key_hash=_meta.get("user_api_key_hash"), user_api_key_hash=_meta.get("user_api_key_hash"),
user_api_key_alias=_meta.get("user_api_key_alias"), user_api_key_alias=_meta.get("user_api_key_alias"),
user_api_key_spend=_meta.get("user_api_key_spend"), user_api_key_spend=_meta.get("user_api_key_spend"),
user_api_key_max_budget=_meta.get("user_api_key_max_budget"), user_api_key_max_budget=_meta.get("user_api_key_max_budget"),
user_api_key_budget_reset_at=_meta.get("user_api_key_budget_reset_at"), user_api_key_budget_reset_at=_meta.get("user_api_key_budget_reset_at"),
user_api_key_org_id=_meta.get("user_api_key_org_id"), user_api_key_org_id=_meta.get("user_api_key_org_id"),
user_api_key_team_id=_meta.get("user_api_key_team_id"), user_api_key_team_id=_meta.get("user_api_key_team_id"),
user_api_key_user_id=_meta.get("user_api_key_user_id"), user_api_key_project_id=_meta.get("user_api_key_project_id"),
user_api_key_team_alias=_meta.get("user_api_key_team_alias"), user_api_key_user_id=_meta.get("user_api_key_user_id"),
user_api_key_end_user_id=_meta.get("user_api_key_end_user_id"), user_api_key_team_alias=_meta.get("user_api_key_team_alias"),
user_api_key_user_email=_meta.get("user_api_key_user_email"), user_api_key_end_user_id=_meta.get("user_api_key_end_user_id"),
user_api_key_request_route=_meta.get("user_api_key_request_route"), user_api_key_user_email=_meta.get("user_api_key_user_email"),
user_api_key_auth_metadata=_meta.get("user_api_key_auth_metadata"), user_api_key_request_route=_meta.get("user_api_key_request_route"),
) user_api_key_auth_metadata=_meta.get("user_api_key_auth_metadata"),
) )
)
# Prune + Possibly alert
window_seconds = self.pagerduty_alerting_args.get( # Prune + Possibly alert
"failure_threshold_window_seconds", 60 window_seconds = self.pagerduty_alerting_args.get(
) "failure_threshold_window_seconds", 60
threshold = self.pagerduty_alerting_args.get("failure_threshold", 1) )
threshold = self.pagerduty_alerting_args.get("failure_threshold", 1)
# If threshold is crossed, send PD alert for failures
await self._send_alert_if_thresholds_crossed( # If threshold is crossed, send PD alert for failures
events=self._failure_events, await self._send_alert_if_thresholds_crossed(
window_seconds=window_seconds, events=self._failure_events,
threshold=threshold, window_seconds=window_seconds,
alert_prefix="High LLM API Failure Rate", threshold=threshold,
) alert_prefix="High LLM API Failure Rate",
)
async def async_pre_call_hook(
self, async def async_pre_call_hook(
user_api_key_dict: UserAPIKeyAuth, self,
cache: DualCache, user_api_key_dict: UserAPIKeyAuth,
data: dict, cache: DualCache,
call_type: CallTypesLiteral, data: dict,
) -> Optional[Union[Exception, str, dict]]: call_type: CallTypesLiteral,
""" ) -> Optional[Union[Exception, str, dict]]:
Example of detecting hanging requests by waiting a given threshold. """
If the request didn't finish by then, we treat it as 'hanging'. Example of detecting hanging requests by waiting a given threshold.
""" If the request didn't finish by then, we treat it as 'hanging'.
verbose_logger.info("Inside Proxy Logging Pre-call hook!") """
asyncio.create_task( verbose_logger.info("Inside Proxy Logging Pre-call hook!")
self.hanging_response_handler( asyncio.create_task(
request_data=data, user_api_key_dict=user_api_key_dict self.hanging_response_handler(
) request_data=data, user_api_key_dict=user_api_key_dict
) )
return None )
return None
async def hanging_response_handler(
self, request_data: Optional[dict], user_api_key_dict: UserAPIKeyAuth async def hanging_response_handler(
): self, request_data: Optional[dict], user_api_key_dict: UserAPIKeyAuth
""" ):
Checks if request completed by the time 'hanging_threshold_seconds' elapses. """
If not, we classify it as a hanging request. Checks if request completed by the time 'hanging_threshold_seconds' elapses.
""" If not, we classify it as a hanging request.
verbose_logger.debug( """
f"Inside Hanging Response Handler!..sleeping for {self.pagerduty_alerting_args.get('hanging_threshold_seconds', PAGERDUTY_DEFAULT_HANGING_THRESHOLD_SECONDS)} seconds" verbose_logger.debug(
) f"Inside Hanging Response Handler!..sleeping for {self.pagerduty_alerting_args.get('hanging_threshold_seconds', PAGERDUTY_DEFAULT_HANGING_THRESHOLD_SECONDS)} seconds"
await asyncio.sleep( )
self.pagerduty_alerting_args.get( await asyncio.sleep(
"hanging_threshold_seconds", PAGERDUTY_DEFAULT_HANGING_THRESHOLD_SECONDS self.pagerduty_alerting_args.get(
) "hanging_threshold_seconds", PAGERDUTY_DEFAULT_HANGING_THRESHOLD_SECONDS
) )
)
if await self._request_is_completed(request_data=request_data):
return # It's not hanging if completed if await self._request_is_completed(request_data=request_data):
return # It's not hanging if completed
# Otherwise, record it as hanging
self._hanging_events.append( # Otherwise, record it as hanging
PagerDutyInternalEvent( self._hanging_events.append(
failure_event_type="hanging_response", PagerDutyInternalEvent(
timestamp=datetime.now(timezone.utc), failure_event_type="hanging_response",
error_class="HangingRequest", timestamp=datetime.now(timezone.utc),
error_code="HangingRequest", error_class="HangingRequest",
error_llm_provider="HangingRequest", error_code="HangingRequest",
user_api_key_hash=user_api_key_dict.api_key, error_llm_provider="HangingRequest",
user_api_key_alias=user_api_key_dict.key_alias, user_api_key_hash=user_api_key_dict.api_key,
user_api_key_spend=user_api_key_dict.spend, user_api_key_alias=user_api_key_dict.key_alias,
user_api_key_max_budget=user_api_key_dict.max_budget, user_api_key_spend=user_api_key_dict.spend,
user_api_key_budget_reset_at=( user_api_key_max_budget=user_api_key_dict.max_budget,
user_api_key_dict.budget_reset_at.isoformat() user_api_key_budget_reset_at=(
if user_api_key_dict.budget_reset_at user_api_key_dict.budget_reset_at.isoformat()
else None if user_api_key_dict.budget_reset_at
), else None
user_api_key_org_id=user_api_key_dict.org_id, ),
user_api_key_team_id=user_api_key_dict.team_id, user_api_key_org_id=user_api_key_dict.org_id,
user_api_key_user_id=user_api_key_dict.user_id, user_api_key_team_id=user_api_key_dict.team_id,
user_api_key_team_alias=user_api_key_dict.team_alias, user_api_key_project_id=user_api_key_dict.project_id,
user_api_key_end_user_id=user_api_key_dict.end_user_id, user_api_key_user_id=user_api_key_dict.user_id,
user_api_key_user_email=user_api_key_dict.user_email, user_api_key_team_alias=user_api_key_dict.team_alias,
user_api_key_request_route=user_api_key_dict.request_route, user_api_key_end_user_id=user_api_key_dict.end_user_id,
user_api_key_auth_metadata=user_api_key_dict.metadata, user_api_key_user_email=user_api_key_dict.user_email,
) user_api_key_request_route=user_api_key_dict.request_route,
) user_api_key_auth_metadata=user_api_key_dict.metadata,
)
# Prune + Possibly alert )
window_seconds = self.pagerduty_alerting_args.get(
"hanging_threshold_window_seconds", # Prune + Possibly alert
PAGERDUTY_DEFAULT_HANGING_THRESHOLD_WINDOW_SECONDS, window_seconds = self.pagerduty_alerting_args.get(
) "hanging_threshold_window_seconds",
threshold: int = self.pagerduty_alerting_args.get( PAGERDUTY_DEFAULT_HANGING_THRESHOLD_WINDOW_SECONDS,
"hanging_threshold_fails", PAGERDUTY_DEFAULT_HANGING_THRESHOLD_SECONDS )
) threshold: int = self.pagerduty_alerting_args.get(
"hanging_threshold_fails", PAGERDUTY_DEFAULT_HANGING_THRESHOLD_SECONDS
# If threshold is crossed, send PD alert for hangs )
await self._send_alert_if_thresholds_crossed(
events=self._hanging_events, # If threshold is crossed, send PD alert for hangs
window_seconds=window_seconds, await self._send_alert_if_thresholds_crossed(
threshold=threshold, events=self._hanging_events,
alert_prefix="High Number of Hanging LLM Requests", window_seconds=window_seconds,
) threshold=threshold,
alert_prefix="High Number of Hanging LLM Requests",
# ------------------ HELPERS ------------------ # )
async def _send_alert_if_thresholds_crossed( # ------------------ HELPERS ------------------ #
self,
events: List[PagerDutyInternalEvent], async def _send_alert_if_thresholds_crossed(
window_seconds: int, self,
threshold: int, events: List[PagerDutyInternalEvent],
alert_prefix: str, window_seconds: int,
): threshold: int,
""" alert_prefix: str,
1. Prune old events ):
2. If threshold is reached, build alert, send to PagerDuty """
3. Clear those events 1. Prune old events
""" 2. If threshold is reached, build alert, send to PagerDuty
cutoff = datetime.now(timezone.utc) - timedelta(seconds=window_seconds) 3. Clear those events
pruned = [e for e in events if e.get("timestamp", datetime.min) > cutoff] """
cutoff = datetime.now(timezone.utc) - timedelta(seconds=window_seconds)
# Update the reference list pruned = [e for e in events if e.get("timestamp", datetime.min) > cutoff]
events.clear()
events.extend(pruned) # Update the reference list
events.clear()
# Check threshold events.extend(pruned)
verbose_logger.debug(
f"Have {len(events)} events in the last {window_seconds} seconds. Threshold is {threshold}" # Check threshold
) verbose_logger.debug(
if len(events) >= threshold: f"Have {len(events)} events in the last {window_seconds} seconds. Threshold is {threshold}"
# Build short summary of last N events )
error_summaries = self._build_error_summaries(events, max_errors=5) if len(events) >= threshold:
alert_message = ( # Build short summary of last N events
f"{alert_prefix}: {len(events)} in the last {window_seconds} seconds." error_summaries = self._build_error_summaries(events, max_errors=5)
) alert_message = (
custom_details = {"recent_errors": error_summaries} f"{alert_prefix}: {len(events)} in the last {window_seconds} seconds."
)
await self.send_alert_to_pagerduty( custom_details = {"recent_errors": error_summaries}
alert_message=alert_message,
custom_details=custom_details, await self.send_alert_to_pagerduty(
) alert_message=alert_message,
custom_details=custom_details,
# Clear them after sending an alert, so we don't spam )
events.clear()
# Clear them after sending an alert, so we don't spam
def _build_error_summaries( events.clear()
self, events: List[PagerDutyInternalEvent], max_errors: int = 5
) -> List[PagerDutyInternalEvent]: def _build_error_summaries(
""" self, events: List[PagerDutyInternalEvent], max_errors: int = 5
Build short text summaries for the last `max_errors`. ) -> List[PagerDutyInternalEvent]:
Example: "ValueError (code: 500, provider: openai)" """
""" Build short text summaries for the last `max_errors`.
recent = events[-max_errors:] Example: "ValueError (code: 500, provider: openai)"
summaries = [] """
for fe in recent: recent = events[-max_errors:]
# If any of these is None, show "N/A" to avoid messing up the summary string summaries = []
fe.pop("timestamp") for fe in recent:
summaries.append(fe) # If any of these is None, show "N/A" to avoid messing up the summary string
return summaries fe.pop("timestamp")
summaries.append(fe)
async def send_alert_to_pagerduty(self, alert_message: str, custom_details: dict): return summaries
"""
Send [critical] Alert to PagerDuty async def send_alert_to_pagerduty(self, alert_message: str, custom_details: dict):
"""
https://developer.pagerduty.com/api-reference/YXBpOjI3NDgyNjU-pager-duty-v2-events-api Send [critical] Alert to PagerDuty
"""
try: https://developer.pagerduty.com/api-reference/YXBpOjI3NDgyNjU-pager-duty-v2-events-api
verbose_logger.debug(f"Sending alert to PagerDuty: {alert_message}") """
async_client: AsyncHTTPHandler = get_async_httpx_client( try:
llm_provider=httpxSpecialProvider.LoggingCallback verbose_logger.debug(f"Sending alert to PagerDuty: {alert_message}")
) async_client: AsyncHTTPHandler = get_async_httpx_client(
payload: PagerDutyRequestBody = PagerDutyRequestBody( llm_provider=httpxSpecialProvider.LoggingCallback
payload=PagerDutyPayload( )
summary=alert_message, payload: PagerDutyRequestBody = PagerDutyRequestBody(
severity="critical", payload=PagerDutyPayload(
source="LiteLLM Alert", summary=alert_message,
component="LiteLLM", severity="critical",
custom_details=custom_details, source="LiteLLM Alert",
), component="LiteLLM",
routing_key=self.api_key, custom_details=custom_details,
event_action="trigger", ),
) routing_key=self.api_key,
event_action="trigger",
return await async_client.post( )
url="https://events.pagerduty.com/v2/enqueue",
json=dict(payload), return await async_client.post(
headers={"Content-Type": "application/json"}, url="https://events.pagerduty.com/v2/enqueue",
) json=dict(payload),
except Exception as e: headers={"Content-Type": "application/json"},
verbose_logger.exception(f"Error sending alert to PagerDuty: {e}") )
except Exception as e:
verbose_logger.exception(f"Error sending alert to PagerDuty: {e}")
File diff suppressed because it is too large Load Diff
+295 -287
View File
@@ -1,287 +1,295 @@
import asyncio import asyncio
import traceback import traceback
from datetime import datetime from datetime import datetime
from typing import Any, List, Optional, Union, cast from typing import Any, List, Optional, Union, cast
import litellm import litellm
from litellm._logging import verbose_proxy_logger from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import ( from litellm.litellm_core_utils.core_helpers import (
_get_parent_otel_span_from_kwargs, _get_parent_otel_span_from_kwargs,
get_litellm_metadata_from_kwargs, get_litellm_metadata_from_kwargs,
) )
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_checks import log_db_metrics from litellm.proxy.auth.auth_checks import log_db_metrics
from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.utils import ProxyUpdateSpend from litellm.proxy.utils import ProxyUpdateSpend
from litellm.types.utils import ( from litellm.types.utils import (
StandardLoggingPayload, StandardLoggingPayload,
StandardLoggingUserAPIKeyMetadata, StandardLoggingUserAPIKeyMetadata,
) )
from litellm.utils import get_end_user_id_for_cost_tracking from litellm.utils import get_end_user_id_for_cost_tracking
class _ProxyDBLogger(CustomLogger): class _ProxyDBLogger(CustomLogger):
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
await self._PROXY_track_cost_callback( await self._PROXY_track_cost_callback(
kwargs, response_obj, start_time, end_time kwargs, response_obj, start_time, end_time
) )
async def async_post_call_failure_hook( async def async_post_call_failure_hook(
self, self,
request_data: dict, request_data: dict,
original_exception: Exception, original_exception: Exception,
user_api_key_dict: UserAPIKeyAuth, user_api_key_dict: UserAPIKeyAuth,
traceback_str: Optional[str] = None, traceback_str: Optional[str] = None,
): ):
request_route = user_api_key_dict.request_route request_route = user_api_key_dict.request_route
if _ProxyDBLogger._should_track_errors_in_db() is False: if _ProxyDBLogger._should_track_errors_in_db() is False:
return return
elif request_route is not None and not RouteChecks.is_llm_api_route( elif request_route is not None and not RouteChecks.is_llm_api_route(
route=request_route route=request_route
): ):
return return
from litellm.proxy.proxy_server import proxy_logging_obj from litellm.proxy.proxy_server import proxy_logging_obj
_metadata = dict( _metadata = dict(
StandardLoggingUserAPIKeyMetadata( StandardLoggingUserAPIKeyMetadata(
user_api_key_hash=user_api_key_dict.api_key, user_api_key_hash=user_api_key_dict.api_key,
user_api_key_alias=user_api_key_dict.key_alias, user_api_key_alias=user_api_key_dict.key_alias,
user_api_key_spend=user_api_key_dict.spend, user_api_key_spend=user_api_key_dict.spend,
user_api_key_max_budget=user_api_key_dict.max_budget, user_api_key_max_budget=user_api_key_dict.max_budget,
user_api_key_budget_reset_at=( user_api_key_budget_reset_at=(
user_api_key_dict.budget_reset_at.isoformat() user_api_key_dict.budget_reset_at.isoformat()
if user_api_key_dict.budget_reset_at if user_api_key_dict.budget_reset_at
else None else None
), ),
user_api_key_user_email=user_api_key_dict.user_email, user_api_key_user_email=user_api_key_dict.user_email,
user_api_key_user_id=user_api_key_dict.user_id, user_api_key_user_id=user_api_key_dict.user_id,
user_api_key_team_id=user_api_key_dict.team_id, user_api_key_team_id=user_api_key_dict.team_id,
user_api_key_org_id=user_api_key_dict.org_id, user_api_key_org_id=user_api_key_dict.org_id,
user_api_key_team_alias=user_api_key_dict.team_alias, user_api_key_project_id=user_api_key_dict.project_id,
user_api_key_end_user_id=user_api_key_dict.end_user_id, user_api_key_team_alias=user_api_key_dict.team_alias,
user_api_key_request_route=user_api_key_dict.request_route, user_api_key_end_user_id=user_api_key_dict.end_user_id,
user_api_key_auth_metadata=user_api_key_dict.metadata, user_api_key_request_route=user_api_key_dict.request_route,
) user_api_key_auth_metadata=user_api_key_dict.metadata,
) )
_metadata["user_api_key"] = user_api_key_dict.api_key )
_metadata["status"] = "failure" _metadata["user_api_key"] = user_api_key_dict.api_key
_metadata["error_information"] = ( _metadata["status"] = "failure"
StandardLoggingPayloadSetup.get_error_information( _metadata[
original_exception=original_exception, "error_information"
traceback_str=traceback_str, ] = StandardLoggingPayloadSetup.get_error_information(
) original_exception=original_exception,
) traceback_str=traceback_str,
)
existing_metadata: dict = request_data.get("metadata", None) or {}
existing_metadata.update(_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"] = {} 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 {} 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"): # Preserve tags from existing metadata
existing_metadata["tags"] = existing_litellm_metadata.get("tags") 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"]["proxy_server_request"] = (
) request_data.get("proxy_server_request")
request_data["litellm_params"]["metadata"] = existing_metadata or existing_litellm_params.get("proxy_server_request")
or {}
# Preserve model name and custom_llm_provider )
if "model" not in request_data: request_data["litellm_params"]["metadata"] = existing_metadata
request_data["model"] = existing_litellm_params.get("model") or request_data.get("model", "")
if "custom_llm_provider" not in request_data: # Preserve model name and custom_llm_provider
request_data["custom_llm_provider"] = existing_litellm_params.get("custom_llm_provider") or request_data.get("custom_llm_provider", "") if "model" not in request_data:
request_data["model"] = existing_litellm_params.get(
await proxy_logging_obj.db_spend_update_writer.update_database( "model"
token=user_api_key_dict.api_key, ) or request_data.get("model", "")
response_cost=0.0, if "custom_llm_provider" not in request_data:
user_id=user_api_key_dict.user_id, request_data["custom_llm_provider"] = existing_litellm_params.get(
end_user_id=user_api_key_dict.end_user_id, "custom_llm_provider"
team_id=user_api_key_dict.team_id, ) or request_data.get("custom_llm_provider", "")
kwargs=request_data,
completion_response=original_exception, await proxy_logging_obj.db_spend_update_writer.update_database(
start_time=datetime.now(), token=user_api_key_dict.api_key,
end_time=datetime.now(), response_cost=0.0,
org_id=user_api_key_dict.org_id, 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,
@log_db_metrics kwargs=request_data,
async def _PROXY_track_cost_callback( completion_response=original_exception,
self, start_time=datetime.now(),
kwargs, # kwargs to completion end_time=datetime.now(),
completion_response: Optional[ org_id=user_api_key_dict.org_id,
Union[litellm.ModelResponse, Any] )
], # response from completion
start_time=None, @log_db_metrics
end_time=None, # start/end time for completion async def _PROXY_track_cost_callback(
): self,
from litellm.proxy.proxy_server import proxy_logging_obj, update_cache kwargs, # kwargs to completion
completion_response: Optional[
verbose_proxy_logger.debug("INSIDE _PROXY_track_cost_callback") Union[litellm.ModelResponse, Any]
try: ], # response from completion
verbose_proxy_logger.debug( start_time=None,
f"kwargs stream: {kwargs.get('stream', None)} + complete streaming response: {kwargs.get('complete_streaming_response', None)}" end_time=None, # start/end time for completion
) ):
parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs=kwargs) from litellm.proxy.proxy_server import proxy_logging_obj, update_cache
litellm_params = kwargs.get("litellm_params", {}) or {}
end_user_id = get_end_user_id_for_cost_tracking(litellm_params) verbose_proxy_logger.debug("INSIDE _PROXY_track_cost_callback")
metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs) try:
user_id = cast(Optional[str], metadata.get("user_api_key_user_id", None)) verbose_proxy_logger.debug(
team_id = cast(Optional[str], metadata.get("user_api_key_team_id", None)) f"kwargs stream: {kwargs.get('stream', None)} + complete streaming response: {kwargs.get('complete_streaming_response', 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)) parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs=kwargs)
end_user_max_budget = metadata.get("user_api_end_user_max_budget", None) litellm_params = kwargs.get("litellm_params", {}) or {}
sl_object: Optional[StandardLoggingPayload] = kwargs.get( end_user_id = get_end_user_id_for_cost_tracking(litellm_params)
"standard_logging_object", None metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs)
) user_id = cast(Optional[str], metadata.get("user_api_key_user_id", None))
response_cost = ( team_id = cast(Optional[str], metadata.get("user_api_key_team_id", None))
sl_object.get("response_cost", None) org_id = cast(Optional[str], metadata.get("user_api_key_org_id", None))
if sl_object is not None key_alias = cast(Optional[str], metadata.get("user_api_key_alias", None))
else kwargs.get("response_cost", None) end_user_max_budget = metadata.get("user_api_end_user_max_budget", None)
) sl_object: Optional[StandardLoggingPayload] = kwargs.get(
tags: Optional[List[str]] = ( "standard_logging_object", None
sl_object.get("request_tags", None) if sl_object is not None else None )
) response_cost = (
sl_object.get("response_cost", None)
if response_cost is not None: if sl_object is not None
user_api_key = metadata.get("user_api_key", None) else kwargs.get("response_cost", None)
if kwargs.get("cache_hit", False) is True: )
response_cost = 0.0 tags: Optional[List[str]] = (
verbose_proxy_logger.debug( sl_object.get("request_tags", None) if sl_object is not None else None
f"Cache Hit: response_cost {response_cost}, for user_id {user_id}" )
)
if response_cost is not None:
verbose_proxy_logger.debug( user_api_key = metadata.get("user_api_key", None)
f"user_api_key {user_api_key}, user_id {user_id}, team_id {team_id}, end_user_id {end_user_id}" if kwargs.get("cache_hit", False) is True:
) response_cost = 0.0
if _should_track_cost_callback( verbose_proxy_logger.debug(
user_api_key=user_api_key, f"Cache Hit: response_cost {response_cost}, for user_id {user_id}"
user_id=user_id, )
team_id=team_id,
end_user_id=end_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}"
## UPDATE DATABASE )
await proxy_logging_obj.db_spend_update_writer.update_database( if _should_track_cost_callback(
token=user_api_key, user_api_key=user_api_key,
response_cost=response_cost, user_id=user_id,
user_id=user_id, team_id=team_id,
end_user_id=end_user_id, end_user_id=end_user_id,
team_id=team_id, ):
kwargs=kwargs, ## UPDATE DATABASE
completion_response=completion_response, await proxy_logging_obj.db_spend_update_writer.update_database(
start_time=start_time, token=user_api_key,
end_time=end_time, response_cost=response_cost,
org_id=org_id, user_id=user_id,
) end_user_id=end_user_id,
team_id=team_id,
# update cache kwargs=kwargs,
asyncio.create_task( completion_response=completion_response,
update_cache( start_time=start_time,
token=user_api_key, end_time=end_time,
user_id=user_id, org_id=org_id,
end_user_id=end_user_id, )
response_cost=response_cost,
team_id=team_id, # update cache
parent_otel_span=parent_otel_span, asyncio.create_task(
tags=tags, update_cache(
) token=user_api_key,
) user_id=user_id,
end_user_id=end_user_id,
await proxy_logging_obj.slack_alerting_instance.customer_spend_alert( response_cost=response_cost,
token=user_api_key, team_id=team_id,
key_alias=key_alias, parent_otel_span=parent_otel_span,
end_user_id=end_user_id, tags=tags,
response_cost=response_cost, )
max_budget=end_user_max_budget, )
)
else: await proxy_logging_obj.slack_alerting_instance.customer_spend_alert(
# Non-model call types (health checks, afile_delete) have no model or standard_logging_object. token=user_api_key,
# Use .get() for "stream" to avoid KeyError on health checks. key_alias=key_alias,
if sl_object is None and not kwargs.get("model"): end_user_id=end_user_id,
verbose_proxy_logger.warning( response_cost=response_cost,
"Cost tracking - skipping, no standard_logging_object and no model for call_type=%s", max_budget=end_user_max_budget,
kwargs.get("call_type", "unknown"), )
) else:
return # Non-model call types (health checks, afile_delete) have no model or standard_logging_object.
if kwargs.get("stream") is not True or ( # Use .get() for "stream" to avoid KeyError on health checks.
kwargs.get("stream") is True and "complete_streaming_response" in kwargs if sl_object is None and not kwargs.get("model"):
): verbose_proxy_logger.warning(
if sl_object is not None: "Cost tracking - skipping, no standard_logging_object and no model for call_type=%s",
cost_tracking_failure_debug_info: Union[dict, str] = ( kwargs.get("call_type", "unknown"),
sl_object["response_cost_failure_debug_info"] # type: ignore )
or "response_cost_failure_debug_info is None in standard_logging_object" return
) if kwargs.get("stream") is not True or (
else: kwargs.get("stream") is True
cost_tracking_failure_debug_info = ( and "complete_streaming_response" in kwargs
"standard_logging_object not found" ):
) if sl_object is not None:
model = kwargs.get("model") cost_tracking_failure_debug_info: Union[dict, str] = (
raise Exception( sl_object["response_cost_failure_debug_info"] # type: ignore
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" or "response_cost_failure_debug_info is None in standard_logging_object"
) )
except Exception as e: else:
error_msg = f"Error in tracking cost callback - {str(e)}\n Traceback:{traceback.format_exc()}" cost_tracking_failure_debug_info = (
model = kwargs.get("model", "") "standard_logging_object not found"
metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs) )
litellm_metadata = kwargs.get("litellm_params", {}).get( model = kwargs.get("model")
"litellm_metadata", {} 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"
old_metadata = kwargs.get("litellm_params", {}).get("metadata", {}) )
call_type = kwargs.get("call_type", "") except Exception as e:
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" error_msg = f"Error in tracking cost callback - {str(e)}\n Traceback:{traceback.format_exc()}"
asyncio.create_task( model = kwargs.get("model", "")
proxy_logging_obj.failed_tracking_alert( metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs)
error_message=error_msg, litellm_metadata = kwargs.get("litellm_params", {}).get(
failing_model=model, "litellm_metadata", {}
) )
) old_metadata = kwargs.get("litellm_params", {}).get("metadata", {})
call_type = kwargs.get("call_type", "")
verbose_proxy_logger.exception( 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"
"Error in tracking cost callback - %s", str(e) asyncio.create_task(
) proxy_logging_obj.failed_tracking_alert(
error_message=error_msg,
@staticmethod failing_model=model,
def _should_track_errors_in_db(): )
""" )
Returns True if errors should be tracked in the database
verbose_proxy_logger.exception(
By default, errors are tracked in the database "Error in tracking cost callback - %s", str(e)
)
If users want to disable error tracking, they can set the disable_error_logs flag in the general_settings
""" @staticmethod
from litellm.proxy.proxy_server import general_settings def _should_track_errors_in_db():
"""
if general_settings.get("disable_error_logs") is True: Returns True if errors should be tracked in the database
return False
return 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
def _should_track_cost_callback( """
user_api_key: Optional[str], from litellm.proxy.proxy_server import general_settings
user_id: Optional[str],
team_id: Optional[str], if general_settings.get("disable_error_logs") is True:
end_user_id: Optional[str], return False
) -> bool: return
"""
Determine if the cost callback should be tracked based on the kwargs
""" def _should_track_cost_callback(
user_api_key: Optional[str],
# don't run track cost callback if user opted into disabling spend user_id: Optional[str],
if ProxyUpdateSpend.disable_spend_updates() is True: team_id: Optional[str],
return False end_user_id: Optional[str],
) -> bool:
if ( """
user_api_key is not None Determine if the cost callback should be tracked based on the kwargs
or user_id is not None """
or team_id is not None
or end_user_id is not None # don't run track cost callback if user opted into disabling spend
): if ProxyUpdateSpend.disable_spend_updates() is True:
return True return False
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
File diff suppressed because it is too large Load Diff