diff --git a/docs/my-website/docs/proxy/call_hooks.md b/docs/my-website/docs/proxy/call_hooks.md index a7b0afcc18ba..c588ca0d0e62 100644 --- a/docs/my-website/docs/proxy/call_hooks.md +++ b/docs/my-website/docs/proxy/call_hooks.md @@ -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 diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md index 4ae82257517e..fdd68c953f60 100644 --- a/docs/my-website/docs/proxy/config_settings.md +++ b/docs/my-website/docs/proxy/config_settings.md @@ -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 diff --git a/litellm/_service_logger.py b/litellm/_service_logger.py index 7a60359d5445..969a9ef14836 100644 --- a/litellm/_service_logger.py +++ b/litellm/_service_logger.py @@ -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 diff --git a/litellm/constants.py b/litellm/constants.py index cf12ec60f070..e224c0dd69a4 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -230,7 +230,7 @@ LITELLM_EMBEDDING_PROVIDERS_SUPPORTING_INPUT_ARRAY_OF_TOKENS = [ "openai", "azure", - "hosted_vllm" + "hosted_vllm", ] @@ -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 ########################### ######################################################################################## diff --git a/litellm/integrations/custom_logger.py b/litellm/integrations/custom_logger.py index 08cebe1a0c10..960dc715e7ef 100644 --- a/litellm/integrations/custom_logger.py +++ b/litellm/integrations/custom_logger.py @@ -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 diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index 921cb161310f..304fa827a7af 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -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 diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 03bf1cd29e8e..aa543ee48911 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -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 diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 012b65581061..88ce34245a60 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -3594,7 +3594,10 @@ 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 "" @@ -3602,14 +3605,14 @@ def get_error_information( _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) diff --git a/litellm/proxy/example_config_yaml/custom_callbacks1.py b/litellm/proxy/example_config_yaml/custom_callbacks1.py index 2cc644a1845c..83f68dd55a24 100644 --- a/litellm/proxy/example_config_yaml/custom_callbacks1.py +++ b/litellm/proxy/example_config_yaml/custom_callbacks1.py @@ -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 diff --git a/litellm/proxy/hooks/parallel_request_limiter_v2.py b/litellm/proxy/hooks/parallel_request_limiter_v2.py index 2c7e024ea84f..8fbd8ad8f125 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v2.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v2.py @@ -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") diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 22e023d9eded..918f3105b03b 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -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: @@ -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 {} diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 2fbedaeb2295..0e7ad685bbf4 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -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 @@ 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 @@ 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 @@ -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, @@ -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 diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 6c551c0b4e80..d1e725def55c 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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. @@ -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 ### @@ -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: diff --git a/tests/litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 43d4dd9cd85d..3aa5240eaea8 100644 --- a/tests/litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -8,7 +8,7 @@ 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 @@ 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 diff --git a/tests/pass_through_unit_tests/test_pass_through_unit_tests.py b/tests/pass_through_unit_tests/test_pass_through_unit_tests.py index dd5fc4275eaa..7f9748343f9c 100644 --- a/tests/pass_through_unit_tests/test_pass_through_unit_tests.py +++ b/tests/pass_through_unit_tests/test_pass_through_unit_tests.py @@ -31,8 +31,8 @@ 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,