Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion docs/my-website/docs/proxy/call_hooks.md
Original file line number Diff line number Diff line change
Expand Up @@ -44,7 +44,8 @@ class MyCustomHandler(CustomLogger): # https://docs.litellm.ai/docs/observabilit
self,
request_data: dict,
original_exception: Exception,
user_api_key_dict: UserAPIKeyAuth
user_api_key_dict: UserAPIKeyAuth,
traceback_str: Optional[str] = None,
):
pass

Expand Down
1 change: 1 addition & 0 deletions docs/my-website/docs/proxy/config_settings.md
Original file line number Diff line number Diff line change
Expand Up @@ -526,6 +526,7 @@ router_settings:
| MAX_TILE_HEIGHT | Maximum height for image tiles. Default is 512
| MAX_TILE_WIDTH | Maximum width for image tiles. Default is 512
| MAX_TOKEN_TRIMMING_ATTEMPTS | Maximum number of attempts to trim a token message. Default is 10
| MAXIMUM_TRACEBACK_LINES_TO_LOG | Maximum number of lines to log in traceback in LiteLLM Logs UI. Default is 100
| MAX_RETRY_DELAY | Maximum delay in seconds for retrying requests. Default is 8.0
| MIN_NON_ZERO_TEMPERATURE | Minimum non-zero temperature value. Default is 0.0001
| MINIMUM_PROMPT_CACHE_TOKEN_COUNT | Minimum token count for caching a prompt. Default is 1024
Expand Down
1 change: 1 addition & 0 deletions litellm/_service_logger.py
Original file line number Diff line number Diff line change
Expand Up @@ -276,6 +276,7 @@ async def async_post_call_failure_hook(
request_data: dict,
original_exception: Exception,
user_api_key_dict: UserAPIKeyAuth,
traceback_str: Optional[str] = None,
):
"""
Hook to track failed litellm-service calls
Expand Down
3 changes: 2 additions & 1 deletion litellm/constants.py
Original file line number Diff line number Diff line change
Expand Up @@ -230,7 +230,7 @@
LITELLM_EMBEDDING_PROVIDERS_SUPPORTING_INPUT_ARRAY_OF_TOKENS = [
"openai",
"azure",
"hosted_vllm"
"hosted_vllm",
]


Expand Down Expand Up @@ -593,6 +593,7 @@
os.getenv("PROMETHEUS_BUDGET_METRICS_REFRESH_INTERVAL_MINUTES", 5)
)
MCP_TOOL_NAME_PREFIX = "mcp_tool"
MAXIMUM_TRACEBACK_LINES_TO_LOG = int(os.getenv("MAXIMUM_TRACEBACK_LINES_TO_LOG", 100))

########################### LiteLLM Proxy Specific Constants ###########################
########################################################################################
Expand Down
1 change: 1 addition & 0 deletions litellm/integrations/custom_logger.py
Original file line number Diff line number Diff line change
Expand Up @@ -234,6 +234,7 @@ async def async_post_call_failure_hook(
request_data: dict,
original_exception: Exception,
user_api_key_dict: UserAPIKeyAuth,
traceback_str: Optional[str] = None,
):
pass

Expand Down
1 change: 1 addition & 0 deletions litellm/integrations/opentelemetry.py
Original file line number Diff line number Diff line change
Expand Up @@ -282,6 +282,7 @@ async def async_post_call_failure_hook(
request_data: dict,
original_exception: Exception,
user_api_key_dict: UserAPIKeyAuth,
traceback_str: Optional[str] = None,
):
from opentelemetry import trace
from opentelemetry.trace import Status, StatusCode
Expand Down
1 change: 1 addition & 0 deletions litellm/integrations/prometheus.py
Original file line number Diff line number Diff line change
Expand Up @@ -802,6 +802,7 @@ async def async_post_call_failure_hook(
request_data: dict,
original_exception: Exception,
user_api_key_dict: UserAPIKeyAuth,
traceback_str: Optional[str] = None,
):
"""
Track client side failures
Expand Down
11 changes: 7 additions & 4 deletions litellm/litellm_core_utils/litellm_logging.py
Original file line number Diff line number Diff line change
Expand Up @@ -3594,22 +3594,25 @@ def strip_trailing_slash(api_base: Optional[str]) -> Optional[str]:
@staticmethod
def get_error_information(
original_exception: Optional[Exception],
traceback_str: Optional[str] = None,
) -> StandardLoggingPayloadErrorInformation:
from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG

error_status: str = str(getattr(original_exception, "status_code", ""))
error_class: str = (
str(original_exception.__class__.__name__) if original_exception else ""
)
_llm_provider_in_exception = getattr(original_exception, "llm_provider", "")

# Get traceback information (first 100 lines)
traceback_info = ""
traceback_info = traceback_str or ""
if original_exception:
tb = getattr(original_exception, "__traceback__", None)
if tb:
import traceback

tb_lines = traceback.format_tb(tb)
traceback_info = "".join(tb_lines[:100]) # Limit to first 100 lines
traceback_info += "".join(
tb_lines[:MAXIMUM_TRACEBACK_LINES_TO_LOG]
) # Limit to first 100 lines

# Get additional error details
error_message = str(original_exception)
Expand Down
1 change: 1 addition & 0 deletions litellm/proxy/example_config_yaml/custom_callbacks1.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,7 @@ async def async_post_call_failure_hook(
request_data: dict,
original_exception: Exception,
user_api_key_dict: UserAPIKeyAuth,
traceback_str: Optional[str] = None,
):
pass

Expand Down
1 change: 1 addition & 0 deletions litellm/proxy/hooks/parallel_request_limiter_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -459,6 +459,7 @@ async def async_post_call_failure_hook(
request_data: dict,
original_exception: Exception,
user_api_key_dict: UserAPIKeyAuth,
traceback_str: Optional[str] = None,
):
try:
self.print_verbose("Inside Max Parallel Request Failure Hook")
Expand Down
2 changes: 2 additions & 0 deletions litellm/proxy/hooks/proxy_track_cost_callback.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@ async def async_post_call_failure_hook(
request_data: dict,
original_exception: Exception,
user_api_key_dict: UserAPIKeyAuth,
traceback_str: Optional[str] = None,
):
request_route = user_api_key_dict.request_route
if _ProxyDBLogger._should_track_errors_in_db() is False:
Expand Down Expand Up @@ -62,6 +63,7 @@ async def async_post_call_failure_hook(
"error_information"
] = StandardLoggingPayloadSetup.get_error_information(
original_exception=original_exception,
traceback_str=traceback_str,
)

existing_metadata: dict = request_data.get("metadata", None) or {}
Expand Down
147 changes: 91 additions & 56 deletions litellm/proxy/pass_through_endpoints/pass_through_endpoints.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import ast
import asyncio
import json
import traceback
import uuid
from base64 import b64encode
from datetime import datetime
Expand All @@ -22,6 +23,7 @@

import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
Expand Down Expand Up @@ -440,10 +442,10 @@ async def make_multipart_http_request(

for field_name, field_value in form_data.items():
if isinstance(field_value, (StarletteUploadFile, UploadFile)):
files[field_name] = (
await HttpPassThroughEndpointHelpers._build_request_files_from_upload_file(
upload_file=field_value
)
files[
field_name
] = await HttpPassThroughEndpointHelpers._build_request_files_from_upload_file(
upload_file=field_value
)
else:
form_data_dict[field_name] = field_value
Expand All @@ -458,6 +460,53 @@ async def make_multipart_http_request(
)
return response

@staticmethod
def _init_kwargs_for_pass_through_endpoint(
request: Request,
user_api_key_dict: UserAPIKeyAuth,
passthrough_logging_payload: PassthroughStandardLoggingPayload,
logging_obj: LiteLLMLoggingObj,
_parsed_body: Optional[dict] = None,
litellm_call_id: Optional[str] = None,
) -> dict:
_parsed_body = _parsed_body or {}
_litellm_metadata: Optional[dict] = _parsed_body.pop("litellm_metadata", None)
_metadata = dict(
StandardLoggingUserAPIKeyMetadata(
user_api_key_hash=user_api_key_dict.api_key,
user_api_key_alias=user_api_key_dict.key_alias,
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_team_id=user_api_key_dict.team_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_end_user_id=user_api_key_dict.end_user_id,
)
)
_metadata["user_api_key"] = user_api_key_dict.api_key
if _litellm_metadata:
_metadata.update(_litellm_metadata)

_metadata = _update_metadata_with_tags_in_header(
request=request,
metadata=_metadata,
)

kwargs = {
"litellm_params": {
"metadata": _metadata,
},
"call_type": "pass_through_endpoint",
"litellm_call_id": litellm_call_id,
"passthrough_logging_payload": passthrough_logging_payload,
}

logging_obj.model_call_details[
"passthrough_logging_payload"
] = passthrough_logging_payload

return kwargs


async def pass_through_request( # noqa: PLR0915
request: Request,
Expand All @@ -470,12 +519,25 @@ async def pass_through_request( # noqa: PLR0915
query_params: Optional[dict] = None,
stream: Optional[bool] = None,
):
"""
Pass through endpoint handler, makes the httpx request for pass-through endpoints and ensures logging hooks are called
"""
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.proxy.proxy_server import proxy_logging_obj

#########################################################
# Initialize variables
#########################################################
litellm_call_id = str(uuid.uuid4())
url: Optional[httpx.URL] = None
try:
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.proxy.proxy_server import proxy_logging_obj

# parsed request body
_parsed_body: Optional[dict] = None
# kwargs for pass through endpoint, contains metadata, litellm_params, call_type, litellm_call_id, passthrough_logging_payload
kwargs: Optional[dict] = None

#########################################################
try:
url = httpx.URL(target)
headers = custom_headers
headers = HttpPassThroughEndpointHelpers.forward_headers_from_request(
Expand All @@ -497,7 +559,6 @@ async def pass_through_request( # noqa: PLR0915
str(url)
)

_parsed_body = None
if custom_body:
_parsed_body = custom_body
else:
Expand Down Expand Up @@ -536,7 +597,7 @@ async def pass_through_request( # noqa: PLR0915
request_body=_parsed_body,
request_method=getattr(request, "method", None),
)
kwargs = _init_kwargs_for_pass_through_endpoint(
kwargs = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint(
user_api_key_dict=user_api_key_dict,
_parsed_body=_parsed_body,
passthrough_logging_payload=passthrough_logging_payload,
Expand Down Expand Up @@ -724,6 +785,27 @@ async def pass_through_request( # noqa: PLR0915
str(e)
)
)

#########################################################
# Monitoring: Trigger post_call_failure_hook
# for pass through endpoint failure
#########################################################
request_payload: dict = _parsed_body or {}
# add user_api_key_dict, litellm_call_id, passthrough_logging_payloa for logging
if kwargs:
for key, value in kwargs.items():
request_payload[key] = value
await proxy_logging_obj.post_call_failure_hook(
user_api_key_dict=user_api_key_dict,
original_exception=e,
request_data=request_payload,
traceback_str=traceback.format_exc(
limit=MAXIMUM_TRACEBACK_LINES_TO_LOG,
),
)

#########################################################

if isinstance(e, HTTPException):
raise ProxyException(
message=getattr(e, "message", str(e.detail)),
Expand All @@ -743,53 +825,6 @@ async def pass_through_request( # noqa: PLR0915
)


def _init_kwargs_for_pass_through_endpoint(
request: Request,
user_api_key_dict: UserAPIKeyAuth,
passthrough_logging_payload: PassthroughStandardLoggingPayload,
logging_obj: LiteLLMLoggingObj,
_parsed_body: Optional[dict] = None,
litellm_call_id: Optional[str] = None,
) -> dict:
_parsed_body = _parsed_body or {}
_litellm_metadata: Optional[dict] = _parsed_body.pop("litellm_metadata", None)
_metadata = dict(
StandardLoggingUserAPIKeyMetadata(
user_api_key_hash=user_api_key_dict.api_key,
user_api_key_alias=user_api_key_dict.key_alias,
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_team_id=user_api_key_dict.team_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_end_user_id=user_api_key_dict.end_user_id,
)
)
_metadata["user_api_key"] = user_api_key_dict.api_key
if _litellm_metadata:
_metadata.update(_litellm_metadata)

_metadata = _update_metadata_with_tags_in_header(
request=request,
metadata=_metadata,
)

kwargs = {
"litellm_params": {
"metadata": _metadata,
},
"call_type": "pass_through_endpoint",
"litellm_call_id": litellm_call_id,
"passthrough_logging_payload": passthrough_logging_payload,
}

logging_obj.model_call_details["passthrough_logging_payload"] = (
passthrough_logging_payload
)

return kwargs


def _update_metadata_with_tags_in_header(request: Request, metadata: dict) -> dict:
"""
If tags are in the request headers, add them to the metadata
Expand Down
10 changes: 10 additions & 0 deletions litellm/proxy/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -778,6 +778,7 @@ async def post_call_failure_hook(
user_api_key_dict: UserAPIKeyAuth,
error_type: Optional[ProxyErrorTypes] = None,
route: Optional[str] = None,
traceback_str: Optional[str] = None,
):
"""
Allows users to raise custom exceptions/log when a call fails, without having to deal with parsing Request body.
Expand All @@ -786,6 +787,14 @@ async def post_call_failure_hook(
1. /chat/completions
2. /embeddings
3. /image/generation

Args:
- request_data: dict - The request data.
- original_exception: Exception - The original exception.
- user_api_key_dict: UserAPIKeyAuth - The user api key dict.
- error_type: Optional[ProxyErrorTypes] - The error type.
- route: Optional[str] - The route.
- traceback_str: Optional[str] - The traceback string, sometimes upstream endpoints might need to send the upstream traceback. In which case we use this
"""

### ALERTING ###
Expand Down Expand Up @@ -840,6 +849,7 @@ async def post_call_failure_hook(
request_data=request_data,
user_api_key_dict=user_api_key_dict,
original_exception=original_exception,
traceback_str=traceback_str,
)
)
except Exception as e:
Expand Down
Loading