[Feat] - Add Support for Showing Passthrough endpoint Error Logs on LiteLLM UI (#10990)

* fix: add error logging for passthrough endpoints

* feat: add error logging for passthrough endpoints

* fix: post_call_failure_hook track errors on pt

* fix: use constant for MAXIMUM_TRACEBACK_LINES_TO_LOG

* docs MAXIMUM_TRACEBACK_LINES_TO_LOG

* test: ensure failure callback triggered

* fix: move _init_kwargs_for_pass_through_endpoint
This commit is contained in:
Ishaan Jaff 2025-05-20 18:29:39 -07:00 • committed by GitHub
parent 98e9db340c
commit 3a6802fef1
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
15 changed files with 190 additions and 67 deletions

View file

@ -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

View file

@ -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

View file

@ -276,6 +276,7 @@ class ServiceLogging(CustomLogger):
request_data: dict,
original_exception: Exception,
user_api_key_dict: UserAPIKeyAuth,
traceback_str: Optional[str] = None,
):
"""
Hook to track failed litellm-service calls

View file

@ -230,7 +230,7 @@ LITELLM_CHAT_PROVIDERS = [
LITELLM_EMBEDDING_PROVIDERS_SUPPORTING_INPUT_ARRAY_OF_TOKENS = [
"openai",
"azure",
"hosted_vllm"
"hosted_vllm",
]
@ -593,6 +593,7 @@ PROMETHEUS_BUDGET_METRICS_REFRESH_INTERVAL_MINUTES = int(
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 ###########################
########################################################################################

View file

@ -234,6 +234,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
request_data: dict,
original_exception: Exception,
user_api_key_dict: UserAPIKeyAuth,
traceback_str: Optional[str] = None,
):
pass

View file

@ -282,6 +282,7 @@ class OpenTelemetry(CustomLogger):
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

View file

@ -802,6 +802,7 @@ class PrometheusLogger(CustomLogger):
request_data: dict,
original_exception: Exception,
user_api_key_dict: UserAPIKeyAuth,
traceback_str: Optional[str] = None,
):
"""
Track client side failures

View file

@ -3594,7 +3594,10 @@ class StandardLoggingPayloadSetup:
@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 ""
@ -3602,14 +3605,14 @@ class StandardLoggingPayloadSetup:
_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)

View file

@ -39,6 +39,7 @@ class MyCustomHandler(
request_data: dict,
original_exception: Exception,
user_api_key_dict: UserAPIKeyAuth,
traceback_str: Optional[str] = None,
):
pass

View file

@ -459,6 +459,7 @@ class _PROXY_MaxParallelRequestsHandler_v2(BaseRoutingStrategy, CustomLogger):
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")

View file

@ -33,6 +33,7 @@ class _ProxyDBLogger(CustomLogger):
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:
@ -62,6 +63,7 @@ class _ProxyDBLogger(CustomLogger):
"error_information"
] = StandardLoggingPayloadSetup.get_error_information(
original_exception=original_exception,
traceback_str=traceback_str,
)
existing_metadata: dict = request_data.get("metadata", None) or {}

View file

@ -1,6 +1,7 @@
import ast
import asyncio
import json
import traceback
import uuid
from base64 import b64encode
from datetime import datetime
@ -22,6 +23,7 @@ from starlette.datastructures import UploadFile as StarletteUploadFile
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
@ -440,10 +442,10 @@ class HttpPassThroughEndpointHelpers:
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
@ -458,6 +460,53 @@ class HttpPassThroughEndpointHelpers:
)
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,
@ -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(
@ -497,7 +559,6 @@ async def pass_through_request( # noqa: PLR0915
str(url)
)
_parsed_body = None
if custom_body:
_parsed_body = custom_body
else:
@ -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,
@ -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)),
@ -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

View file

@ -778,6 +778,7 @@ class ProxyLogging:
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.
@ -786,6 +787,14 @@ class ProxyLogging:
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 ###
@ -840,6 +849,7 @@ class ProxyLogging:
request_data=request_data,
user_api_key_dict=user_api_key_dict,
original_exception=original_exception,
traceback_str=traceback_str,
)
)
except Exception as e:

View file

@ -8,7 +8,7 @@ import httpx
import pytest
from fastapi import Request, UploadFile
from fastapi.testclient import TestClient
from starlette.datastructures import Headers
from starlette.datastructures import Headers, QueryParams
from starlette.datastructures import UploadFile as StarletteUploadFile
sys.path.insert(
@ -17,6 +17,7 @@ sys.path.insert(
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
HttpPassThroughEndpointHelpers,
pass_through_request,
)
@ -114,3 +115,66 @@ async def test_make_multipart_http_request():
assert isinstance(call_args["files"], dict)
assert isinstance(call_args["data"], dict)
assert call_args["data"]["text_field"] == "test value"
@pytest.mark.asyncio
async def test_pass_through_request_failure_handler():
"""
Test that the failure handler is called when pass_through_request fails
Critical Test: When a users pass through endpoint request fails, we must log the failure code, exception in litellm spend logs.
"""
print("running test_pass_through_request_failure_handler")
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
with patch(
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client"
) as mock_get_client:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.ProxyBaseLLMRequestProcessing"
) as mock_processing:
# Setup mock for post_call_failure_hook and pre_call_hook
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.pre_call_hook = AsyncMock()
# Setup mock for httpx client
mock_client = MagicMock()
mock_client.client = MagicMock()
mock_client.client.request = AsyncMock(
side_effect=httpx.HTTPError("Request failed")
)
mock_get_client.return_value = mock_client
# Mock headers for custom headers
mock_processing.get_custom_headers.return_value = {}
# Create mock request
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.body = AsyncMock(return_value=b'{"test": "data"}')
mock_request.headers = Headers({})
# Create a simple empty QueryParams
mock_request.query_params = QueryParams({})
# Create mock user API key dict
mock_user_api_key_dict = MagicMock()
# Call the function with a target that will trigger an HTTPError
with pytest.raises(Exception):
await pass_through_request(
request=mock_request,
target="http://test.com",
custom_headers={},
user_api_key_dict=mock_user_api_key_dict,
)
# Assert post_call_failure_hook was called
mock_proxy_logging.post_call_failure_hook.assert_called_once()
# Verify the arguments to post_call_failure_hook
call_args = mock_proxy_logging.post_call_failure_hook.call_args[1]
assert call_args["user_api_key_dict"] == mock_user_api_key_dict
assert isinstance(
call_args["original_exception"], TypeError
) # Now expecting TypeError
assert "traceback_str" in call_args

View file

@ -31,8 +31,8 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
from fastapi import Request
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
_init_kwargs_for_pass_through_endpoint,
_update_metadata_with_tags_in_header,
HttpPassThroughEndpointHelpers
)
from litellm.types.passthrough_endpoints.pass_through_endpoints import PassthroughStandardLoggingPayload
@ -110,7 +110,7 @@ def test_init_kwargs_for_pass_through_endpoint_basic(
request_body={},
)
result = _init_kwargs_for_pass_through_endpoint(
result = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint(
request=request,
user_api_key_dict=mock_user_api_key_dict,
passthrough_logging_payload=passthrough_payload,
@ -161,7 +161,7 @@ def test_init_kwargs_with_litellm_metadata(mock_request, mock_user_api_key_dict)
request_body={},
)
result = _init_kwargs_for_pass_through_endpoint(
result = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint(
request=request,
user_api_key_dict=mock_user_api_key_dict,
passthrough_logging_payload=passthrough_payload,
@ -196,7 +196,7 @@ def test_init_kwargs_with_tags_in_header(mock_request, mock_user_api_key_dict):
request_body={},
)
result = _init_kwargs_for_pass_through_endpoint(
result = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint(
request=request,
user_api_key_dict=mock_user_api_key_dict,
passthrough_logging_payload=passthrough_payload,