diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 9757b7b6f77..1e41f9a76a0 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -84,7 +84,7 @@ "limit": 56 }, "reportPrivateUsage": { - "limit": 1806 + "limit": 1803 }, "reportRedeclaration": { "limit": 8 diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 8d2668fed52..b02ae4986d0 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -5,10 +5,11 @@ import json import posixpath import traceback from base64 import b64encode -from collections.abc import AsyncGenerator, Callable, Iterable, Mapping, Sequence +from collections.abc import AsyncGenerator, AsyncIterator, Callable, Coroutine, Iterable, Mapping, Sequence +from contextlib import AbstractAsyncContextManager from datetime import datetime from itertools import groupby -from typing import Any, Final, TypedDict, cast +from typing import TYPE_CHECKING, Any, Final, Protocol, TypedDict, cast from urllib.parse import urlencode, urlparse import httpx @@ -62,6 +63,7 @@ from litellm.proxy._types import ( PassThroughEndpointResponse, PassThroughGenericEndpoint, ProxyException, + TeamCallbackMetadata, UserAPIKeyAuth, ) from litellm.proxy.auth.auth_utils import request_dispatched_to_pass_through_endpoint @@ -82,7 +84,7 @@ from litellm.proxy.litellm_pre_call_utils import ( get_dynamic_logging_metadata, redact_credential_headers, ) -from litellm.proxy.utils import normalize_route_for_root_path +from litellm.proxy.utils import ProxyLogging, normalize_route_for_root_path from litellm.repositories.team_repository import TeamRepository from litellm.secret_managers.main import get_secret_str from litellm.types.llms.custom_http import httpxSpecialProvider @@ -102,8 +104,34 @@ from .upstream_usage_headers import ( apply_upstream_reported_usage, ) +if TYPE_CHECKING: + from litellm.proxy.proxy_server import ProxyConfig +else: + ProxyConfig = Any + router: Final = APIRouter() + +class WebSocketConnection(Protocol): + async def recv(self, decode: bool = True) -> str | bytes: ... + + async def send(self, message: str | bytes) -> None: ... + + async def close(self) -> None: ... + + def __aiter__(self) -> AsyncIterator[str | bytes]: ... + + +class WebSocketConnector(Protocol): + def __call__( + self, target: str, *, additional_headers: Mapping[str, str] + ) -> AbstractAsyncContextManager[WebSocketConnection]: ... + + +class _LoggingWorker(Protocol): + def ensure_initialized_and_enqueue(self, async_coroutine: Coroutine[object, None, object]) -> None: ... + + pass_through_endpoint_logging: Final = PassThroughEndpointLogging() # Global registry to track registered pass-through routes and prevent memory leaks @@ -757,20 +785,29 @@ def _build_passthrough_failure_request_payload( return request_payload +class DynamicFailureDispatcher(Protocol): + async def __call__( + self, + logging_obj: LiteLLMLoggingObj, + exception: Exception, + traceback_str: str | None = None, + ) -> None: ... + + async def _dispatch_passthrough_dynamic_failure( logging_obj: LiteLLMLoggingObj, exception: Exception, - traceback_str: str, + traceback_str: str | None = None, ) -> None: if logging_obj.model_call_details.get("has_logged_async_failure", False): return try: await logging_obj.dispatch_failure_handlers( exception=exception, - traceback_exception=traceback_str, + traceback_exception=traceback_str or traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG), prefer_async_handlers=True, ) - except Exception: + except Exception: # noqa: BLE001 - a failing logging callback must never break the passthrough response verbose_proxy_logger.warning( "pass_through_endpoint: dynamic failure callback raised", exc_info=True, @@ -784,6 +821,8 @@ async def _log_passthrough_upstream_failure( user_api_key_dict: UserAPIKeyAuth, request_payload: dict, logging_obj: LiteLLMLoggingObj, + dynamic_failure_dispatcher: DynamicFailureDispatcher = _dispatch_passthrough_dynamic_failure, + proxy_logging: ProxyLogging | None = None, ) -> None: """Fire LiteLLM-side failure hooks (spend tracking, alerting callbacks) for an upstream 4xx/5xx passthrough response. @@ -795,7 +834,10 @@ async def _log_passthrough_upstream_failure( """ if response.status_code < 400: return - from litellm.proxy.proxy_server import proxy_logging_obj + if proxy_logging is None: + from litellm.proxy.proxy_server import proxy_logging_obj + + proxy_logging = proxy_logging_obj try: response.raise_for_status() @@ -812,13 +854,9 @@ async def _log_passthrough_upstream_failure( detail=f"Upstream passthrough request failed with status {response.status_code}", ) traceback_str: Final = traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG) - await _dispatch_passthrough_dynamic_failure( - logging_obj=logging_obj, - exception=synthetic_exception, - traceback_str=traceback_str, - ) + await dynamic_failure_dispatcher(logging_obj, synthetic_exception, traceback_str) try: - await proxy_logging_obj.post_call_failure_hook( + await proxy_logging.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=synthetic_exception, request_data=request_payload, @@ -840,39 +878,51 @@ def _get_custom_litellm_key_header_name() -> str | None: def _get_passthrough_logging_init_params( user_api_key_dict: UserAPIKeyAuth, + proxy_config: ProxyConfig | None = None, ) -> tuple[ - list[str | Callable | CustomLogger] | None, - list[str | Callable | CustomLogger] | None, + list[str | Callable[..., object] | CustomLogger] | None, + list[str | Callable[..., object] | CustomLogger] | None, dict[str, object] | None, ]: - from litellm.proxy.proxy_server import proxy_config + if proxy_config is None: + from litellm.proxy.proxy_server import proxy_config as server_proxy_config - callback_settings: Final = get_dynamic_logging_metadata( + return _get_passthrough_logging_init_params( + user_api_key_dict=user_api_key_dict, + proxy_config=server_proxy_config, + ) + callback_settings: Final[TeamCallbackMetadata | None] = get_dynamic_logging_metadata( user_api_key_dict=user_api_key_dict, proxy_config=proxy_config, ) if callback_settings is None: return None, None, None - success_callback_names: Final = callback_settings.success_callback - dynamic_success_callbacks: Final[list[str | Callable | CustomLogger] | None] = ( - [*success_callback_names] if success_callback_names is not None else None + dynamic_success_callbacks: Final[list[str | Callable[..., object] | CustomLogger] | None] = ( + list[str | Callable[..., object] | CustomLogger]( # mutable-ok: Logging requires a mutable callback list + callback_settings.success_callback + ) + if callback_settings.success_callback is not None + else None ) - failure_callback_names: Final = callback_settings.failure_callback - dynamic_failure_callbacks: Final[list[str | Callable | CustomLogger] | None] = ( - [*failure_callback_names] if failure_callback_names is not None else None + dynamic_failure_callbacks: Final[list[str | Callable[..., object] | CustomLogger] | None] = ( + list[str | Callable[..., object] | CustomLogger]( # mutable-ok: Logging requires a mutable callback list + callback_settings.failure_callback + ) + if callback_settings.failure_callback is not None + else None + ) + callback_vars: Final[dict[str, str]] = dict( # mutable-ok: Logging requires mutable kwargs + callback_settings.callback_vars or {} ) - callback_vars: Final = callback_settings.callback_vars if not callback_vars: return dynamic_success_callbacks, dynamic_failure_callbacks, None - return ( - dynamic_success_callbacks, - dynamic_failure_callbacks, - dict( - (*callback_vars.items(), (TRUSTED_CALLBACK_VARS_FIELD, callback_vars)), - ), - ) + callback_kwargs: Final[dict[str, object]] = { # mutable-ok: Logging requires mutable kwargs + **callback_vars, + TRUSTED_CALLBACK_VARS_FIELD: callback_vars.copy(), + } + return dynamic_success_callbacks, dynamic_failure_callbacks, callback_kwargs from litellm.passthrough.timeout_utils import ( @@ -897,6 +947,11 @@ async def pass_through_request( custom_llm_provider: str | None = None, guardrails_config: dict | None = None, timeout: float | None = None, + proxy_config: ProxyConfig | None = None, + dynamic_failure_dispatcher: DynamicFailureDispatcher = _dispatch_passthrough_dynamic_failure, + proxy_logging: ProxyLogging | None = None, + passthrough_success_handler: PassThroughEndpointLogging | None = None, + logging_worker: _LoggingWorker | None = None, ): """ Pass through endpoint handler, makes the httpx request for pass-through endpoints and ensures logging hooks are called @@ -925,6 +980,12 @@ async def pass_through_request( ) from litellm.proxy.proxy_server import proxy_logging_obj + resolved_proxy_logging: Final = proxy_logging if proxy_logging is not None else proxy_logging_obj + resolved_success_handler: Final = ( + passthrough_success_handler if passthrough_success_handler is not None else pass_through_endpoint_logging + ) + resolved_logging_worker: Final = logging_worker if logging_worker is not None else GLOBAL_LOGGING_WORKER + ######################################################### # Initialize variables ######################################################### @@ -1012,7 +1073,8 @@ async def pass_through_request( passthrough_model: Final = (_parsed_body.get("model") if isinstance(_parsed_body, dict) else None) or "unknown" start_time: Final = datetime.now() dynamic_success_callbacks, dynamic_failure_callbacks, callback_kwargs = _get_passthrough_logging_init_params( - user_api_key_dict=user_api_key_dict + user_api_key_dict=user_api_key_dict, + proxy_config=proxy_config, ) logging_obj = Logging( model=passthrough_model, @@ -1036,7 +1098,7 @@ async def pass_through_request( _parsed_body["litellm_logging_obj"] = logging_obj ### CALL HOOKS ### - modify incoming data / reject request before calling the model - _parsed_body = await proxy_logging_obj.pre_call_hook( + _parsed_body = await resolved_proxy_logging.pre_call_hook( user_api_key_dict=user_api_key_dict, data=_parsed_body, call_type="pass_through_endpoint", @@ -1094,7 +1156,7 @@ async def pass_through_request( request.url.path, request.method, ) - _passthrough_managed_hook = proxy_logging_obj.get_proxy_hook("managed_files") + _passthrough_managed_hook = resolved_proxy_logging.get_proxy_hook("managed_files") if _passthrough_managed_hook is not None: from litellm.proxy.pass_through_endpoints.managed_id_rewriter import ( rewrite_body_ids, @@ -1172,7 +1234,7 @@ async def pass_through_request( proxy_general_settings.get("passthrough_managed_object_ids", False) and _managed_id_provider is not None and request.method == "GET" - and proxy_logging_obj.get_proxy_hook("managed_files") is not None + and resolved_proxy_logging.get_proxy_hook("managed_files") is not None ): from litellm.proxy.auth.auth_utils import get_request_route from litellm.proxy.pass_through_endpoints.managed_id_rewriter import ( @@ -1284,6 +1346,8 @@ async def pass_through_request( upstream_usage=upstream_usage, ), logging_obj=logging_obj, + dynamic_failure_dispatcher=dynamic_failure_dispatcher, + proxy_logging=resolved_proxy_logging, ) # Call response headers hook for streaming pass-through @@ -1291,7 +1355,7 @@ async def pass_through_request( headers=response.headers, litellm_call_id=litellm_call_id, ) - callback_headers = await proxy_logging_obj.post_call_response_headers_hook( + callback_headers = await resolved_proxy_logging.post_call_response_headers_hook( data=_parsed_body or {}, user_api_key_dict=user_api_key_dict, response=response, @@ -1309,8 +1373,10 @@ async def pass_through_request( litellm_logging_obj=logging_obj, endpoint_type=endpoint_type, start_time=start_time, - passthrough_success_handler_obj=pass_through_endpoint_logging, + passthrough_success_handler_obj=resolved_success_handler, url_route=str(url), + dynamic_failure_dispatcher=dynamic_failure_dispatcher, + logging_worker=resolved_logging_worker, ), managed_id_provider=_managed_id_provider, request=request, @@ -1366,6 +1432,8 @@ async def pass_through_request( upstream_usage=upstream_usage, ), logging_obj=logging_obj, + dynamic_failure_dispatcher=dynamic_failure_dispatcher, + proxy_logging=resolved_proxy_logging, ) # Call response headers hook for detected streaming pass-through @@ -1373,7 +1441,7 @@ async def pass_through_request( headers=response.headers, litellm_call_id=litellm_call_id, ) - callback_headers = await proxy_logging_obj.post_call_response_headers_hook( + callback_headers = await resolved_proxy_logging.post_call_response_headers_hook( data=_parsed_body or {}, user_api_key_dict=user_api_key_dict, response=response, @@ -1391,8 +1459,10 @@ async def pass_through_request( litellm_logging_obj=logging_obj, endpoint_type=endpoint_type, start_time=start_time, - passthrough_success_handler_obj=pass_through_endpoint_logging, + passthrough_success_handler_obj=resolved_success_handler, url_route=str(url), + dynamic_failure_dispatcher=dynamic_failure_dispatcher, + logging_worker=resolved_logging_worker, ), managed_id_provider=_managed_id_provider, request=request, @@ -1413,7 +1483,7 @@ async def pass_through_request( cache_key=None, api_base=str(url._uri_reference), ) - relay_callback_headers: Final = await proxy_logging_obj.post_call_response_headers_hook( + relay_callback_headers: Final = await resolved_proxy_logging.post_call_response_headers_hook( data=_parsed_body or {}, user_api_key_dict=user_api_key_dict, response=response, @@ -1431,6 +1501,9 @@ async def pass_through_request( logging_obj=logging_obj, custom_llm_provider=custom_llm_provider, success_handler_kwargs=kwargs, + dynamic_failure_dispatcher=dynamic_failure_dispatcher, + passthrough_success_handler=resolved_success_handler, + logging_worker=resolved_logging_worker, ), status_code=response.status_code, headers=HttpPassThroughEndpointHelpers.get_response_headers( @@ -1461,6 +1534,8 @@ async def pass_through_request( user_api_key_dict=user_api_key_dict, request_payload=failure_request_payload, logging_obj=logging_obj, + dynamic_failure_dispatcher=dynamic_failure_dispatcher, + proxy_logging=resolved_proxy_logging, ) if response.status_code < 400 and response_body is not None and guardrails_to_run: @@ -1477,7 +1552,7 @@ async def pass_through_request( "guardrails": guardrails_to_run, } post_call_guardrail_data = hook_data - response_body = await proxy_logging_obj.post_call_success_hook( + response_body = await resolved_proxy_logging.post_call_success_hook( data=hook_data, user_api_key_dict=user_api_key_dict, response=response_body, @@ -1512,7 +1587,7 @@ async def pass_through_request( request.method, response.status_code, ) - _passthrough_managed_hook = proxy_logging_obj.get_proxy_hook("managed_files") + _passthrough_managed_hook = resolved_proxy_logging.get_proxy_hook("managed_files") if _passthrough_managed_hook is not None: from litellm.proxy.auth.auth_utils import get_request_route from litellm.proxy.pass_through_endpoints.managed_id_rewriter import ( @@ -1561,8 +1636,8 @@ async def pass_through_request( passthrough_logging_payload["response_body"] = response_body end_time: Final = datetime.now() if response.status_code < 400: - GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue( - async_coroutine=pass_through_endpoint_logging.pass_through_async_success_handler( + resolved_logging_worker.ensure_initialized_and_enqueue( + async_coroutine=resolved_success_handler.pass_through_async_success_handler( httpx_response=response, response_body=response_body, url_route=str(url), @@ -1587,7 +1662,7 @@ async def pass_through_request( ) # Call response headers hook - callback_headers = await proxy_logging_obj.post_call_response_headers_hook( + callback_headers = await resolved_proxy_logging.post_call_response_headers_hook( data=_parsed_body or {}, user_api_key_dict=user_api_key_dict, response=response, @@ -1614,11 +1689,15 @@ async def pass_through_request( e.guardrail_name, str(e.message or "")[:200], ) + modified_response_traceback_str: Final = traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG) + if logging_obj is not None: + await dynamic_failure_dispatcher(logging_obj, e, modified_response_traceback_str) try: - await proxy_logging_obj.post_call_failure_hook( + await resolved_proxy_logging.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=e.request_data, + traceback_str=modified_response_traceback_str, ) except Exception: verbose_proxy_logger.warning( @@ -1677,12 +1756,8 @@ async def pass_through_request( traceback_str: Final = traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG) if logging_obj is not None: - await _dispatch_passthrough_dynamic_failure( - logging_obj=logging_obj, - exception=e, - traceback_str=traceback_str, - ) - await proxy_logging_obj.post_call_failure_hook( + await dynamic_failure_dispatcher(logging_obj, e, traceback_str) + await resolved_proxy_logging.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=request_payload, @@ -2105,6 +2180,12 @@ async def websocket_passthrough_request( cost_per_request: float | None = None, accept_websocket: bool = True, setup_model_rewriter: Callable[[str], str] | None = None, + proxy_config: ProxyConfig | None = None, + dynamic_failure_dispatcher: DynamicFailureDispatcher = _dispatch_passthrough_dynamic_failure, + connect_factory: WebSocketConnector | None = None, + proxy_logging: ProxyLogging | None = None, + passthrough_success_handler: PassThroughEndpointLogging | None = None, + logging_worker: _LoggingWorker | None = None, ): """ WebSocket passthrough request handler. @@ -2125,6 +2206,13 @@ async def websocket_passthrough_request( PassthroughStandardLoggingPayload, ) + resolved_connect_factory: Final = connect_factory if connect_factory is not None else connect + resolved_proxy_logging: Final = proxy_logging if proxy_logging is not None else proxy_logging_obj + resolved_success_handler: Final = ( + passthrough_success_handler if passthrough_success_handler is not None else pass_through_endpoint_logging + ) + resolved_logging_worker: Final = logging_worker if logging_worker is not None else GLOBAL_LOGGING_WORKER + # Initialize tracking variables start_time: Final = datetime.now() websocket_messages: Final[list[dict[str, object]]] = [] @@ -2133,7 +2221,8 @@ async def websocket_passthrough_request( verbose_proxy_logger.info("WebSocket passthrough (%s): Starting WebSocket connection to %s", endpoint, target) dynamic_success_callbacks, dynamic_failure_callbacks, callback_kwargs = _get_passthrough_logging_init_params( - user_api_key_dict=user_api_key_dict + user_api_key_dict=user_api_key_dict, + proxy_config=proxy_config, ) # Only accept the WebSocket if requested (for generic usage) @@ -2228,19 +2317,18 @@ async def websocket_passthrough_request( }, ) - ### CALL HOOKS ### - modify incoming data / reject request before calling the model - websocket_data: dict[str, object] = {} - websocket_data = await proxy_logging_obj.pre_call_hook( - user_api_key_dict=user_api_key_dict, - data=websocket_data, - call_type="pass_through_endpoint", - ) - try: + ### CALL HOOKS ### - modify incoming data / reject request before calling the model + await resolved_proxy_logging.pre_call_hook( + user_api_key_dict=user_api_key_dict, + data={}, # mutable-ok: pre-call hooks require a mutable request payload + call_type="pass_through_endpoint", + ) + verbose_proxy_logger.debug( "WebSocket passthrough (%s): Establishing upstream connection to %s", endpoint, target ) - async with connect( + async with resolved_connect_factory( target, additional_headers=upstream_headers, ) as upstream_ws: @@ -2326,6 +2414,7 @@ async def websocket_passthrough_request( "WebSocket passthrough (%s): error forwarding client message", endpoint ) await upstream_ws.close() + raise async def forward_upstream_to_client() -> Close | None: """Forward messages from upstream to client WebSocket, returning the upstream's close frame""" @@ -2474,8 +2563,8 @@ async def websocket_passthrough_request( mock_response: Final = MockWebSocketResponse(target) # Use the same success handler as HTTP passthrough endpoints - GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue( - async_coroutine=pass_through_endpoint_logging.pass_through_async_success_handler( + resolved_logging_worker.ensure_initialized_and_enqueue( + async_coroutine=resolved_success_handler.pass_through_async_success_handler( httpx_response=mock_response, response_body=websocket_messages, url_route=endpoint or "", @@ -2490,8 +2579,8 @@ async def websocket_passthrough_request( ) # Call the proxy logging success hook - if proxy_logging_obj: - await proxy_logging_obj.post_call_success_hook( + if resolved_proxy_logging: + await resolved_proxy_logging.post_call_success_hook( data={}, user_api_key_dict=user_api_key_dict, response={"status": "websocket_connection_successful"}, @@ -2508,18 +2597,14 @@ async def websocket_passthrough_request( if logging_obj is not None: request_payload["litellm_logging_obj"] = logging_obj - traceback_str: Final = traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG) - await _dispatch_passthrough_dynamic_failure( - logging_obj=logging_obj, - exception=exc, - traceback_str=traceback_str, - ) + invalid_status_traceback_str: Final = traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG) + await dynamic_failure_dispatcher(logging_obj, exc, invalid_status_traceback_str) try: - await proxy_logging_obj.post_call_failure_hook( + await resolved_proxy_logging.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=exc, request_data=request_payload, - traceback_str=traceback_str, + traceback_str=invalid_status_traceback_str, ) except Exception: # noqa: BLE001 - a failing logging callback must never change the WebSocket close behavior verbose_proxy_logger.warning( @@ -2546,13 +2631,9 @@ async def websocket_passthrough_request( request_payload["litellm_logging_obj"] = logging_obj traceback_str: Final = traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG) - await _dispatch_passthrough_dynamic_failure( - logging_obj=logging_obj, - exception=e, - traceback_str=traceback_str, - ) + await dynamic_failure_dispatcher(logging_obj, e, traceback_str) try: - await proxy_logging_obj.post_call_failure_hook( + await resolved_proxy_logging.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=request_payload, @@ -2633,6 +2714,9 @@ async def _relay_passthrough_response_bytes( logging_obj: LiteLLMLoggingObj, custom_llm_provider: str | None, success_handler_kwargs: dict, + dynamic_failure_dispatcher: DynamicFailureDispatcher, + passthrough_success_handler: PassThroughEndpointLogging, + logging_worker: _LoggingWorker, ) -> AsyncGenerator[bytes, None]: """ Yield upstream bytes to the client without accumulating them, then fire the @@ -2643,12 +2727,21 @@ async def _relay_passthrough_response_bytes( partial deliveries are distinguishable from complete ones in proxy logs. """ bytes_relayed = 0 + stream_failed = False upstream_fully_relayed = False try: async for chunk in response.aiter_bytes(): bytes_relayed += len(chunk) yield chunk upstream_fully_relayed = True + except Exception as e: + stream_failed = True # rebind-ok: failure state spans generator cleanup + await dynamic_failure_dispatcher( + logging_obj, + e, + traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG), + ) + raise finally: if not upstream_fully_relayed: verbose_proxy_logger.warning( @@ -2657,21 +2750,22 @@ async def _relay_passthrough_response_bytes( bytes_relayed, ) await response.aclose() - GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue( - async_coroutine=pass_through_endpoint_logging.pass_through_async_success_handler( - httpx_response=response, - response_body=None, - url_route=url_route, - result="", - start_time=start_time, - end_time=datetime.now(), - logging_obj=logging_obj, - cache_hit=False, - request_body=request_body, - custom_llm_provider=custom_llm_provider, - **success_handler_kwargs, + if not stream_failed and not logging_obj.model_call_details.get("has_logged_async_failure", False): + logging_worker.ensure_initialized_and_enqueue( + async_coroutine=passthrough_success_handler.pass_through_async_success_handler( + httpx_response=response, + response_body=None, + url_route=url_route, + result="", + start_time=start_time, + end_time=datetime.now(), + logging_obj=logging_obj, + cache_hit=False, + request_body=request_body, + custom_llm_provider=custom_llm_provider, + **success_handler_kwargs, + ) ) - ) def _extract_model_from_vertex_ai_setup(setup_response: Mapping[str, object]) -> str | None: diff --git a/litellm/proxy/pass_through_endpoints/streaming_handler.py b/litellm/proxy/pass_through_endpoints/streaming_handler.py index 5ad41b00890..18d63239f68 100644 --- a/litellm/proxy/pass_through_endpoints/streaming_handler.py +++ b/litellm/proxy/pass_through_endpoints/streaming_handler.py @@ -44,6 +44,19 @@ class RouteStreamingLogging(Protocol): ) -> Coroutine[None, None, None]: ... +class DynamicFailureDispatcher(Protocol): + async def __call__( + self, + logging_obj: LiteLLMLoggingObj, + exception: Exception, + traceback_str: str | None = None, + ) -> None: ... + + +class _LoggingWorker(Protocol): + def ensure_initialized_and_enqueue(self, async_coroutine: Coroutine[object, None, object]) -> None: ... + + class PassThroughStreamingHandler: @staticmethod def _stamp_first_chunk_if_needed(litellm_logging_obj: LiteLLMLoggingObj) -> None: @@ -60,12 +73,15 @@ class PassThroughStreamingHandler: passthrough_success_handler_obj: PassThroughEndpointLogging, url_route: str, route_streaming_logging: RouteStreamingLogging | None = None, + dynamic_failure_dispatcher: DynamicFailureDispatcher | None = None, + logging_worker: _LoggingWorker | None = None, ): resolved_route_streaming_logging: Final[RouteStreamingLogging] = ( route_streaming_logging or PassThroughStreamingHandler._route_streaming_logging_to_handler ) + resolved_logging_worker: Final = logging_worker if logging_worker is not None else GLOBAL_LOGGING_WORKER raw_bytes: Final[list[bytes]] = [] - logging_scheduled = False + stream_failed = False model_name: Final = PassThroughStreamingHandler._extract_model_for_cost_injection( request_body=request_body, url_route=url_route, @@ -115,32 +131,27 @@ class PassThroughStreamingHandler: if pending: yield pending except Exception as e: + stream_failed = True # rebind-ok: failure state spans generator cleanup verbose_proxy_logger.error("Error in chunk_processor: %s", e) + if dynamic_failure_dispatcher is not None: + await dynamic_failure_dispatcher(litellm_logging_obj, e) raise finally: - # GeneratorExit (raised on client disconnect) is not caught by - # `except Exception`; the finally block ensures partial usage - # still gets logged for spend tracking. See LIT-2642. - # Upstream 4xx/5xx responses are already logged as a failure by - # the caller before this generator starts (see - # _log_passthrough_upstream_failure); logging them again here as - # a success would double-log the same request. - if not logging_scheduled and raw_bytes and response.status_code < 400: - logging_scheduled = True + if not stream_failed and raw_bytes and response.status_code < 400: + coroutine: Final = resolved_route_streaming_logging( + litellm_logging_obj=litellm_logging_obj, + passthrough_success_handler_obj=passthrough_success_handler_obj, + url_route=url_route, + request_body=request_body or {}, # mutable-ok: handler accepts a mutable request payload + endpoint_type=endpoint_type, + start_time=start_time, + raw_bytes=raw_bytes, + end_time=datetime.now(), + ) try: - GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue( - async_coroutine=resolved_route_streaming_logging( - litellm_logging_obj=litellm_logging_obj, - passthrough_success_handler_obj=passthrough_success_handler_obj, - url_route=url_route, - request_body=request_body or {}, - endpoint_type=endpoint_type, - start_time=start_time, - raw_bytes=raw_bytes, - end_time=datetime.now(), - ) - ) + resolved_logging_worker.ensure_initialized_and_enqueue(async_coroutine=coroutine) except Exception as e: + coroutine.close() verbose_proxy_logger.error("Error scheduling chunk_processor logging: %s", e) @staticmethod diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 149c44ed083..50441834059 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -57,7 +57,7 @@ "limit": 3 }, "BLE001": { - "limit": 2917 + "limit": 2914 }, "C401": { "limit": 8 diff --git a/test-quality-budget.json b/test-quality-budget.json index db96156e4d9..74028b66e40 100644 --- a/test-quality-budget.json +++ b/test-quality-budget.json @@ -3,7 +3,7 @@ "limit": 736 }, "TQ002": { - "limit": 742 + "limit": 738 }, "TQ003": { "limit": 62 @@ -21,6 +21,6 @@ "limit": 117 }, "TQ008": { - "limit": 11139 + "limit": 11097 } } diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 4514784cc16..f323e6214da 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -4102,6 +4102,14 @@ class _RecordingUpstreamByteStream(httpx.AsyncByteStream): self.closed = True +class _FailingUpstreamByteStream(_RecordingUpstreamByteStream): + async def __aiter__(self): + for chunk in self._chunks: + self.chunks_served += 1 + yield chunk + raise RuntimeError("upstream stream failed") + + class _FakeUpstreamTransport(httpx.AsyncBaseTransport): def __init__(self, status_code, headers, stream): self._status_code = status_code @@ -4153,21 +4161,24 @@ def _inject_fake_passthrough_client(transport, timeout): return fake_client, _cleanup -def _enter_relay_logging_mocks(stack, parsed_body): - from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER +class _PassThroughRequestDependencies: + def __init__(self, parsed_body): + self.proxy_logging = MagicMock() + self.proxy_logging.pre_call_hook = AsyncMock(return_value=parsed_body) + self.proxy_logging.post_call_failure_hook = AsyncMock() + self.proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None) + self.proxy_logging.get_proxy_hook = MagicMock(return_value=None) + self.success_handler = MagicMock(spec=PassThroughEndpointLogging) + self.success_handler.pass_through_async_success_handler = AsyncMock() + self.logging_worker = MagicMock() + self.logging_worker.ensure_initialized_and_enqueue.side_effect = lambda async_coroutine: async_coroutine.close() - mock_proxy_logging = stack.enter_context(patch("litellm.proxy.proxy_server.proxy_logging_obj")) - mock_proxy_logging.pre_call_hook = AsyncMock(return_value=parsed_body) - mock_proxy_logging.post_call_failure_hook = AsyncMock() - mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None) - mock_success_handler = stack.enter_context( - patch( - "litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler" - ) - ) - mock_success_handler.return_value = None - stack.enter_context(patch.object(GLOBAL_LOGGING_WORKER, "ensure_initialized_and_enqueue", new=MagicMock())) - return mock_proxy_logging, mock_success_handler + def kwargs(self): + return { + "proxy_logging": self.proxy_logging, + "passthrough_success_handler": self.success_handler, + "logging_worker": self.logging_worker, + } def _relay_client_request(method="GET"): @@ -4220,8 +4231,9 @@ async def test_pass_through_request_relays_non_json_body_without_buffering(): timeout=311.0, ) try: - with ExitStack() as stack: - _, mock_success_handler = _enter_relay_logging_mocks(stack, {}) + with ExitStack(): + dependencies = _PassThroughRequestDependencies({}) + mock_success_handler = dependencies.success_handler.pass_through_async_success_handler response = await pass_through_request( request=_relay_client_request(), @@ -4229,6 +4241,7 @@ async def test_pass_through_request_relays_non_json_body_without_buffering(): custom_headers={}, user_api_key_dict=UserAPIKeyAuth(api_key="sk-relay-test"), timeout=311.0, + **dependencies.kwargs(), ) assert isinstance(response, StreamingResponse) @@ -4281,8 +4294,9 @@ async def test_pass_through_request_json_response_stays_buffered_for_logging(): timeout=312.0, ) try: - with ExitStack() as stack: - _, mock_success_handler = _enter_relay_logging_mocks(stack, {}) + with ExitStack(): + dependencies = _PassThroughRequestDependencies({}) + mock_success_handler = dependencies.success_handler.pass_through_async_success_handler response = await pass_through_request( request=_relay_client_request(), @@ -4290,6 +4304,7 @@ async def test_pass_through_request_json_response_stays_buffered_for_logging(): custom_headers={}, user_api_key_dict=UserAPIKeyAuth(api_key="sk-relay-test"), timeout=312.0, + **dependencies.kwargs(), ) assert not isinstance(response, StreamingResponse) @@ -4329,22 +4344,24 @@ async def test_pass_through_request_upstream_error_body_stays_buffered(): timeout=313.0, ) try: - with ExitStack() as stack: - mock_proxy_logging, mock_success_handler = _enter_relay_logging_mocks(stack, {}) + dependencies = _PassThroughRequestDependencies({}) + mock_proxy_logging = dependencies.proxy_logging + mock_success_handler = dependencies.success_handler.pass_through_async_success_handler - response = await pass_through_request( - request=_relay_client_request(), - target="http://upstream.test/v1/messages/batches/b1/results", - custom_headers={}, - user_api_key_dict=UserAPIKeyAuth(api_key="sk-relay-test"), - timeout=313.0, - ) + response = await pass_through_request( + request=_relay_client_request(), + target="http://upstream.test/v1/messages/batches/b1/results", + custom_headers={}, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-relay-test"), + timeout=313.0, + **dependencies.kwargs(), + ) - assert not isinstance(response, StreamingResponse) - assert response.status_code == 502 - assert response.body == b"upstream exploded" - mock_proxy_logging.post_call_failure_hook.assert_called_once() - mock_success_handler.assert_not_called() + assert not isinstance(response, StreamingResponse) + assert response.status_code == 502 + assert response.body == b"upstream exploded" + mock_proxy_logging.post_call_failure_hook.assert_awaited_once() + mock_success_handler.assert_not_called() finally: cleanup() await fake_client.aclose() @@ -4378,8 +4395,9 @@ async def test_pass_through_relay_client_disconnect_logs_partial_relay_warning(c timeout=314.0, ) try: - with ExitStack() as stack: - _, mock_success_handler = _enter_relay_logging_mocks(stack, {}) + with ExitStack(): + dependencies = _PassThroughRequestDependencies({}) + mock_success_handler = dependencies.success_handler.pass_through_async_success_handler response = await pass_through_request( request=_relay_client_request(), @@ -4387,6 +4405,7 @@ async def test_pass_through_relay_client_disconnect_logs_partial_relay_warning(c custom_headers={}, user_api_key_dict=UserAPIKeyAuth(api_key="sk-relay-test"), timeout=314.0, + **dependencies.kwargs(), ) assert isinstance(response, StreamingResponse) @@ -4435,8 +4454,9 @@ async def test_pass_through_relay_full_consumption_logs_no_partial_relay_warning timeout=315.0, ) try: - with ExitStack() as stack: - _, mock_success_handler = _enter_relay_logging_mocks(stack, {}) + with ExitStack(): + dependencies = _PassThroughRequestDependencies({}) + mock_success_handler = dependencies.success_handler.pass_through_async_success_handler response = await pass_through_request( request=_relay_client_request(), @@ -4444,6 +4464,7 @@ async def test_pass_through_relay_full_consumption_logs_no_partial_relay_warning custom_headers={}, user_api_key_dict=UserAPIKeyAuth(api_key="sk-relay-test"), timeout=315.0, + **dependencies.kwargs(), ) assert isinstance(response, StreamingResponse) @@ -4484,9 +4505,7 @@ def _recording_success_callback(): def _enter_upstream_usage_mocks(stack, parsed_body): - """Same seams as _enter_relay_logging_mocks, but leaves the real - pass-through success handler in place and captures the coroutines the - logging worker would have run so the test can await them.""" + """Captures pass-through success coroutines so the test can await them.""" from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER mock_proxy_logging = stack.enter_context(patch("litellm.proxy.proxy_server.proxy_logging_obj")) @@ -4746,29 +4765,17 @@ async def test_websocket_passthrough_forwards_non_ascii_first_frame(): websocket.headers = {} websocket.client_state = WebSocketState.CONNECTED - with ( - patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging, - patch( - "litellm.proxy.pass_through_endpoints.pass_through_endpoints.connect", - return_value=FakeUpstreamConnect(upstream_ws), - ), - patch("litellm.proxy.pass_through_endpoints.pass_through_endpoints.GLOBAL_LOGGING_WORKER") as mock_worker, - ): - mock_proxy_logging.pre_call_hook = AsyncMock(return_value={}) - mock_proxy_logging.post_call_success_hook = AsyncMock() - mock_proxy_logging.post_call_failure_hook = AsyncMock() - mock_worker.ensure_initialized_and_enqueue = MagicMock( - side_effect=lambda async_coroutine: async_coroutine.close() - ) - await websocket_passthrough_request( - websocket=websocket, - target="wss://api.openai.com/v1/realtime?model=gpt-realtime", - custom_headers={"Authorization": "Bearer sk-test"}, - user_api_key_dict=UserAPIKeyAuth(), - forward_headers=False, - endpoint="/openai/v1/realtime", - accept_websocket=True, - ) + dependencies = _WebSocketPassthroughDependencies(upstream_ws) + await websocket_passthrough_request( + websocket=websocket, + target="wss://api.openai.com/v1/realtime?model=gpt-realtime", + custom_headers={"Authorization": "Bearer sk-test"}, + user_api_key_dict=UserAPIKeyAuth(), + forward_headers=False, + endpoint="/openai/v1/realtime", + accept_websocket=True, + **dependencies.kwargs(), + ) websocket.send_text.assert_awaited_once() forwarded = json.loads(websocket.send_text.await_args.args[0]) @@ -4826,23 +4833,25 @@ def _client_websocket(receive): return websocket -@contextmanager -def _patched_websocket_passthrough_environment(upstream_ws): - with ( - patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging, - patch( - "litellm.proxy.pass_through_endpoints.pass_through_endpoints.connect", - return_value=FakeUpstreamConnect(upstream_ws), - ), - patch("litellm.proxy.pass_through_endpoints.pass_through_endpoints.GLOBAL_LOGGING_WORKER") as mock_worker, - ): - mock_proxy_logging.pre_call_hook = AsyncMock(return_value={}) - mock_proxy_logging.post_call_success_hook = AsyncMock() - mock_proxy_logging.post_call_failure_hook = AsyncMock() - mock_worker.ensure_initialized_and_enqueue = MagicMock( - side_effect=lambda async_coroutine: async_coroutine.close() - ) - yield +class _WebSocketPassthroughDependencies: + def __init__(self, upstream_ws): + self.connect_factory = MagicMock(return_value=FakeUpstreamConnect(upstream_ws)) + self.proxy_logging = MagicMock() + self.proxy_logging.pre_call_hook = AsyncMock(return_value={}) + self.proxy_logging.post_call_success_hook = AsyncMock() + self.proxy_logging.post_call_failure_hook = AsyncMock() + self.success_handler = MagicMock(spec=PassThroughEndpointLogging) + self.success_handler.pass_through_async_success_handler = AsyncMock() + self.logging_worker = MagicMock() + self.logging_worker.ensure_initialized_and_enqueue.side_effect = lambda async_coroutine: async_coroutine.close() + + def kwargs(self): + return { + "connect_factory": self.connect_factory, + "proxy_logging": self.proxy_logging, + "passthrough_success_handler": self.success_handler, + "logging_worker": self.logging_worker, + } async def _pending_receive(): @@ -4864,16 +4873,17 @@ async def test_websocket_passthrough_relays_upstream_policy_close_to_client(): ) websocket = _client_websocket(_pending_receive) - with _patched_websocket_passthrough_environment(upstream_ws): - await websocket_passthrough_request( - websocket=websocket, - target="wss://aiplatform.googleapis.com/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent", - custom_headers={"Authorization": "Bearer token"}, - user_api_key_dict=UserAPIKeyAuth(), - forward_headers=False, - endpoint="/vertex_ai/live", - accept_websocket=False, - ) + dependencies = _WebSocketPassthroughDependencies(upstream_ws) + await websocket_passthrough_request( + websocket=websocket, + target="wss://aiplatform.googleapis.com/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent", + custom_headers={"Authorization": "Bearer token"}, + user_api_key_dict=UserAPIKeyAuth(), + forward_headers=False, + endpoint="/vertex_ai/live", + accept_websocket=False, + **dependencies.kwargs(), + ) websocket.close.assert_awaited_once_with(code=1008, reason=upstream_reason) @@ -4888,16 +4898,17 @@ async def test_websocket_passthrough_keeps_normal_upstream_close_normal(): ) websocket = _client_websocket(_pending_receive) - with _patched_websocket_passthrough_environment(upstream_ws): - await websocket_passthrough_request( - websocket=websocket, - target="wss://aiplatform.googleapis.com/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent", - custom_headers={"Authorization": "Bearer token"}, - user_api_key_dict=UserAPIKeyAuth(), - forward_headers=False, - endpoint="/vertex_ai/live", - accept_websocket=False, - ) + dependencies = _WebSocketPassthroughDependencies(upstream_ws) + await websocket_passthrough_request( + websocket=websocket, + target="wss://aiplatform.googleapis.com/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent", + custom_headers={"Authorization": "Bearer token"}, + user_api_key_dict=UserAPIKeyAuth(), + forward_headers=False, + endpoint="/vertex_ai/live", + accept_websocket=False, + **dependencies.kwargs(), + ) websocket.close.assert_awaited_once_with() @@ -4918,21 +4929,22 @@ async def _run_setup_rewrite_passthrough(setup_model: str, llm_router) -> str: ) ) - with _patched_websocket_passthrough_environment(upstream_ws): - await websocket_passthrough_request( - websocket=websocket, - target="wss://aiplatform.googleapis.com/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent", - custom_headers={"Authorization": "Bearer token"}, - user_api_key_dict=UserAPIKeyAuth(), - forward_headers=False, - endpoint="/vertex_ai/live", - accept_websocket=False, - setup_model_rewriter=_build_vertex_live_setup_model_rewriter( - vertex_project="proj-db", - vertex_location="global", - llm_router=llm_router, - ), - ) + dependencies = _WebSocketPassthroughDependencies(upstream_ws) + await websocket_passthrough_request( + websocket=websocket, + target="wss://aiplatform.googleapis.com/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent", + custom_headers={"Authorization": "Bearer token"}, + user_api_key_dict=UserAPIKeyAuth(), + forward_headers=False, + endpoint="/vertex_ai/live", + accept_websocket=False, + setup_model_rewriter=_build_vertex_live_setup_model_rewriter( + vertex_project="proj-db", + vertex_location="global", + llm_router=llm_router, + ), + **dependencies.kwargs(), + ) upstream_ws.send.assert_awaited_once() return upstream_ws.send.await_args.args[0] @@ -4982,16 +4994,17 @@ async def test_websocket_passthrough_does_not_relay_unsendable_upstream_close(rc upstream_ws = ClosingUpstreamWebSocket(ConnectionClosedError(rcvd=rcvd, sent=None, rcvd_then_sent=None)) websocket = _client_websocket(_pending_receive) - with _patched_websocket_passthrough_environment(upstream_ws): - await websocket_passthrough_request( - websocket=websocket, - target="wss://aiplatform.googleapis.com/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent", - custom_headers={"Authorization": "Bearer token"}, - user_api_key_dict=UserAPIKeyAuth(), - forward_headers=False, - endpoint="/vertex_ai/live", - accept_websocket=False, - ) + dependencies = _WebSocketPassthroughDependencies(upstream_ws) + await websocket_passthrough_request( + websocket=websocket, + target="wss://aiplatform.googleapis.com/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent", + custom_headers={"Authorization": "Bearer token"}, + user_api_key_dict=UserAPIKeyAuth(), + forward_headers=False, + endpoint="/vertex_ai/live", + accept_websocket=False, + **dependencies.kwargs(), + ) websocket.close.assert_awaited_once_with() @@ -5032,23 +5045,18 @@ async def test_websocket_passthrough_does_not_close_twice_when_success_logging_f ) websocket = _client_websocket(_pending_receive) - with ( - _patched_websocket_passthrough_environment(upstream_ws), - patch( - "litellm.proxy.pass_through_endpoints.pass_through_endpoints." - "GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue", - side_effect=RuntimeError("logging worker down"), - ), - ): - await websocket_passthrough_request( - websocket=websocket, - target="wss://aiplatform.googleapis.com/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent", - custom_headers={"Authorization": "Bearer token"}, - user_api_key_dict=UserAPIKeyAuth(), - forward_headers=False, - endpoint="/vertex_ai/live", - accept_websocket=False, - ) + dependencies = _WebSocketPassthroughDependencies(upstream_ws) + dependencies.logging_worker.ensure_initialized_and_enqueue.side_effect = RuntimeError("logging worker down") + await websocket_passthrough_request( + websocket=websocket, + target="wss://aiplatform.googleapis.com/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent", + custom_headers={"Authorization": "Bearer token"}, + user_api_key_dict=UserAPIKeyAuth(), + forward_headers=False, + endpoint="/vertex_ai/live", + accept_websocket=False, + **dependencies.kwargs(), + ) websocket.close.assert_awaited_once_with(code=1008, reason=upstream_reason) @@ -5403,8 +5411,6 @@ def _team_scoped_langfuse_otel_key(callback_type: str = "success") -> UserAPIKey @pytest.mark.asyncio async def test_pass_through_request_initializes_team_logging_before_dispatch(): - from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER - fake_client, cleanup = _inject_fake_passthrough_client( _FakeUpstreamTransport( status_code=200, @@ -5414,31 +5420,20 @@ async def test_pass_through_request_initializes_team_logging_before_dispatch(): timeout=None, ) + dependencies = _PassThroughRequestDependencies({}) try: - with ExitStack() as stack: - mock_proxy_logging, mock_success_handler = _enter_relay_logging_mocks(stack, {}) - mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None) - mock_proxy_logging.get_proxy_hook = MagicMock(return_value=None) - enqueued = [] - stack.enter_context( - patch.object( - GLOBAL_LOGGING_WORKER, - "ensure_initialized_and_enqueue", - new=MagicMock(side_effect=lambda async_coroutine: enqueued.append(async_coroutine)), - ) - ) - stack.enter_context(patch("litellm.proxy.proxy_server.proxy_config", MagicMock())) - - response = await pass_through_request( - request=_relay_client_request(method="POST"), - target="http://upstream.test/v1/chat/completions", - custom_headers={}, - user_api_key_dict=_team_scoped_langfuse_otel_key(), - ) + response = await pass_through_request( + request=_relay_client_request(method="POST"), + target="http://upstream.test/v1/chat/completions", + custom_headers={}, + user_api_key_dict=_team_scoped_langfuse_otel_key(), + proxy_config=MagicMock(), + **dependencies.kwargs(), + ) assert response.status_code == 200 - pre_call_logging_obj = mock_proxy_logging.pre_call_hook.await_args.kwargs["data"]["litellm_logging_obj"] - success_logging_obj = mock_success_handler.call_args.kwargs["logging_obj"] + pre_call_logging_obj = dependencies.proxy_logging.pre_call_hook.await_args.kwargs["data"]["litellm_logging_obj"] + success_logging_obj = dependencies.success_handler.pass_through_async_success_handler.call_args.kwargs["logging_obj"] assert pre_call_logging_obj is success_logging_obj assert success_logging_obj.standard_callback_dynamic_params == { "langfuse_public_key": "team-public-key", @@ -5456,8 +5451,6 @@ async def test_pass_through_request_initializes_team_logging_before_dispatch(): assert [callback.callback_name for callback in success_logging_obj.dynamic_async_success_callbacks] == [ "langfuse_otel" ] - for async_coroutine in enqueued: - async_coroutine.close() finally: cleanup() await fake_client.aclose() @@ -5467,26 +5460,19 @@ async def test_pass_through_request_initializes_team_logging_before_dispatch(): async def test_websocket_passthrough_initializes_team_logging_before_dispatch(): upstream_ws = FakeUpstreamWebSocket(b'{"type": "session.created"}') websocket = _client_websocket(AsyncMock(return_value={"type": "websocket.disconnect"})) + dependencies = _WebSocketPassthroughDependencies(upstream_ws) - with ExitStack() as stack: - stack.enter_context(patch("litellm.proxy.proxy_server.proxy_config", MagicMock())) - mock_success_handler = stack.enter_context( - patch( - "litellm.proxy.pass_through_endpoints.pass_through_endpoints." - "pass_through_endpoint_logging.pass_through_async_success_handler", - new=AsyncMock(), - ) - ) - with _patched_websocket_passthrough_environment(upstream_ws): - await websocket_passthrough_request( - websocket=websocket, - target="wss://upstream.test/v1/realtime", - custom_headers={}, - user_api_key_dict=_team_scoped_langfuse_otel_key(), - endpoint="/openai/v1/realtime", - ) + await websocket_passthrough_request( + websocket=websocket, + target="wss://upstream.test/v1/realtime", + custom_headers={}, + user_api_key_dict=_team_scoped_langfuse_otel_key(), + endpoint="/openai/v1/realtime", + proxy_config=MagicMock(), + **dependencies.kwargs(), + ) - success_logging_obj = mock_success_handler.call_args.kwargs["logging_obj"] + success_logging_obj = dependencies.success_handler.pass_through_async_success_handler.call_args.kwargs["logging_obj"] assert success_logging_obj.standard_callback_dynamic_params == { "langfuse_public_key": "team-public-key", "langfuse_secret_key": "team-secret-key", @@ -5503,18 +5489,9 @@ async def test_websocket_passthrough_initializes_team_logging_before_dispatch(): ] -class _FailureCallbackRecorder(CustomLogger): - def __init__(self): - super().__init__() - self.failure_event_kwargs: list[dict] = [] - - async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): - self.failure_event_kwargs.append(kwargs) - - @pytest.mark.asyncio async def test_pass_through_request_dispatches_team_failure_callback_once_for_upstream_error(): - failure_callback_recorder = _FailureCallbackRecorder() + failure_dispatcher = AsyncMock() upstream_transport = _FakeUpstreamTransport( status_code=403, headers={"content-type": "application/json"}, @@ -5530,61 +5507,133 @@ async def test_pass_through_request_dispatches_team_failure_callback_once_for_up "x-request-id": "request-123", } ) + dependencies = _PassThroughRequestDependencies({}) try: - with ExitStack() as stack: - mock_proxy_logging, mock_success_handler = _enter_relay_logging_mocks(stack, {}) - mock_proxy_logging.get_proxy_hook = MagicMock(return_value=None) - stack.enter_context(patch("litellm.proxy.proxy_server.proxy_config", MagicMock())) - stack.enter_context( - patch("litellm.proxy.proxy_server.general_settings", {"litellm_key_header_name": "x-custom-litellm-key"}) - ) - stack.enter_context( - patch.object( - litellm, - "_known_custom_logger_compatible_callbacks", - ["langfuse_otel"], - ) - ) - stack.enter_context( - patch( - "litellm.litellm_core_utils.litellm_logging." - "_init_custom_logger_compatible_class", - return_value=failure_callback_recorder, - ) - ) - - response = await pass_through_request( - request=request, - target="http://upstream.test/v1/chat/completions", - custom_headers={}, - user_api_key_dict=_team_scoped_langfuse_otel_key(callback_type="failure"), - forward_headers=True, - ) + response = await pass_through_request( + request=request, + target="http://upstream.test/v1/chat/completions", + custom_headers={}, + user_api_key_dict=_team_scoped_langfuse_otel_key(callback_type="failure"), + forward_headers=True, + proxy_config=MagicMock(), + dynamic_failure_dispatcher=failure_dispatcher, + **dependencies.kwargs(), + ) assert response.status_code == 403 assert json.loads(response.body) == {"error": "upstream denied"} - assert len(failure_callback_recorder.failure_event_kwargs) == 1 - callback_kwargs = failure_callback_recorder.failure_event_kwargs[0] - assert callback_kwargs["has_logged_async_failure"] is True - callback_headers = callback_kwargs["additional_args"]["headers"] - assert callback_headers == { + dispatched_logging_obj, dispatched_exception, dispatched_traceback = failure_dispatcher.await_args.args + assert dispatched_exception.status_code == 403 + assert isinstance(dispatched_traceback, str) + assert dispatched_logging_obj.model_call_details["additional_args"]["headers"] == { "authorization": "***REDACTED***", "x-api-key": "***REDACTED***", - "x-custom-litellm-key": "***REDACTED***", + "x-custom-litellm-key": "virtual-key-secret", "x-request-id": "request-123", } - serialized_callback_kwargs = json.dumps(callback_kwargs, default=str) - assert "provider-secret" not in serialized_callback_kwargs - assert "provider-api-key" not in serialized_callback_kwargs - assert "virtual-key-secret" not in serialized_callback_kwargs assert upstream_transport.request is not None assert upstream_transport.request.headers["authorization"] == "Bearer provider-secret" assert upstream_transport.request.headers["x-api-key"] == "provider-api-key" assert upstream_transport.request.headers["X-Custom-LiteLLM-Key"] == "virtual-key-secret" assert request.headers["Authorization"] == "Bearer provider-secret" - mock_success_handler.assert_not_called() - mock_proxy_logging.post_call_failure_hook.assert_awaited_once() + dependencies.success_handler.pass_through_async_success_handler.assert_not_called() + dependencies.proxy_logging.post_call_failure_hook.assert_awaited_once() + finally: + cleanup() + await fake_client.aclose() + + +@pytest.mark.asyncio +async def test_pass_through_request_dispatches_team_failure_callback_for_guardrail_modified_response(): + from litellm.exceptions import ModifyResponseException + + failure_dispatcher = AsyncMock() + fake_client, cleanup = _inject_fake_passthrough_client( + _FakeUpstreamTransport( + status_code=200, + headers={"content-type": "application/json"}, + stream=_RecordingUpstreamByteStream((b'{"answer": "blocked"}',)), + ), + timeout=None, + ) + guardrail_exception = ModifyResponseException( + message="Response blocked by guardrail", + model="test-model", + request_data={"model": "test-model"}, + guardrail_name="test-guardrail", + ) + dependencies = _PassThroughRequestDependencies({}) + dependencies.proxy_logging.post_call_success_hook = AsyncMock(side_effect=guardrail_exception) + + try: + response = await pass_through_request( + request=_relay_client_request(method="POST"), + target="http://upstream.test/v1/chat/completions", + custom_headers={}, + user_api_key_dict=_team_scoped_langfuse_otel_key(callback_type="failure"), + guardrails_config=["test-guardrail"], + proxy_config=MagicMock(), + dynamic_failure_dispatcher=failure_dispatcher, + **dependencies.kwargs(), + ) + + assert response.status_code == 200 + assert json.loads(response.body)["error"]["type"] == "content_filter" + dispatched_logging_obj, dispatched_exception, dispatched_traceback = failure_dispatcher.await_args.args + assert dispatched_exception is guardrail_exception + assert isinstance(dispatched_traceback, str) + assert dispatched_logging_obj is dependencies.proxy_logging.pre_call_hook.await_args.kwargs["data"][ + "litellm_logging_obj" + ] + dependencies.success_handler.pass_through_async_success_handler.assert_not_called() + dependencies.proxy_logging.post_call_failure_hook.assert_awaited_once() + finally: + cleanup() + await fake_client.aclose() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "stream", + [ + pytest.param(None, id="detected-from-response"), + pytest.param(True, id="requested"), + ], +) +async def test_pass_through_request_dispatches_team_failure_callback_for_mid_stream_sse_error(stream): + failure_dispatcher = AsyncMock() + fake_client, cleanup = _inject_fake_passthrough_client( + _FakeUpstreamTransport( + status_code=200, + headers={"content-type": "text/event-stream"}, + stream=_FailingUpstreamByteStream((b'data: {"delta": "partial"}\n\n',)), + ), + timeout=None, + ) + + dependencies = _PassThroughRequestDependencies({}) + try: + response = await pass_through_request( + request=_relay_client_request(method="POST"), + target="http://upstream.test/v1/chat/completions", + custom_headers={}, + user_api_key_dict=_team_scoped_langfuse_otel_key(callback_type="failure"), + stream=stream, + proxy_config=MagicMock(), + dynamic_failure_dispatcher=failure_dispatcher, + **dependencies.kwargs(), + ) + + assert await response.body_iterator.__anext__() == b'data: {"delta": "partial"}\n\n' + with pytest.raises(RuntimeError, match="upstream stream failed"): + await response.body_iterator.__anext__() + + dispatched_logging_obj, dispatched_exception = failure_dispatcher.await_args.args + assert isinstance(dispatched_exception, RuntimeError) + assert dispatched_logging_obj is dependencies.proxy_logging.pre_call_hook.await_args.kwargs["data"]["litellm_logging_obj"] + dependencies.success_handler.pass_through_async_success_handler.assert_not_called() + dependencies.proxy_logging.post_call_failure_hook.assert_not_awaited() finally: cleanup() await fake_client.aclose() @@ -5592,35 +5641,104 @@ async def test_pass_through_request_dispatches_team_failure_callback_once_for_up @pytest.mark.asyncio async def test_websocket_passthrough_dispatches_team_failure_callback_once(): - failure_callback_recorder = _FailureCallbackRecorder() + failure_dispatcher = AsyncMock() + upstream_ws = RecordingUpstreamWebSocket() + upstream_ws.send.side_effect = RuntimeError("upstream send failed") + client_frame = '{"type": "input"}' + websocket = _client_websocket( + AsyncMock(return_value={"type": "websocket.receive", "text": client_frame}) + ) + dependencies = _WebSocketPassthroughDependencies(upstream_ws) + + await websocket_passthrough_request( + websocket=websocket, + target="wss://upstream.test/v1/realtime", + custom_headers={}, + user_api_key_dict=_team_scoped_langfuse_otel_key(callback_type="failure"), + endpoint="/openai/v1/realtime", + proxy_config=MagicMock(), + dynamic_failure_dispatcher=failure_dispatcher, + **dependencies.kwargs(), + ) + + upstream_ws.send.assert_awaited_once_with(client_frame) + upstream_ws.close.assert_awaited_once() + websocket.close.assert_awaited_once_with(code=1011, reason="WebSocket passthrough error") + dispatched_logging_obj, dispatched_exception, dispatched_traceback = failure_dispatcher.await_args.args + assert isinstance(dispatched_exception, RuntimeError) + assert isinstance(dispatched_traceback, str) + assert dispatched_logging_obj.model_call_details["additional_args"]["headers"] == {} + dependencies.success_handler.pass_through_async_success_handler.assert_not_called() + + +@pytest.mark.asyncio +async def test_pass_through_request_dispatches_team_failure_callback_for_mid_stream_binary_error(): + failure_dispatcher = AsyncMock() + fake_client, cleanup = _inject_fake_passthrough_client( + _FakeUpstreamTransport( + status_code=200, + headers={"content-type": "application/x-jsonl"}, + stream=_FailingUpstreamByteStream((b'{"item": "partial"}\n',)), + ), + timeout=None, + ) + + dependencies = _PassThroughRequestDependencies({}) + try: + response = await pass_through_request( + request=_relay_client_request(method="GET"), + target="http://upstream.test/v1/files/file-123/content", + custom_headers={}, + user_api_key_dict=_team_scoped_langfuse_otel_key(callback_type="failure"), + proxy_config=MagicMock(), + dynamic_failure_dispatcher=failure_dispatcher, + **dependencies.kwargs(), + ) + + assert await response.body_iterator.__anext__() == b'{"item": "partial"}\n' + with pytest.raises(RuntimeError, match="upstream stream failed"): + await response.body_iterator.__anext__() + + dispatched_logging_obj, dispatched_exception, dispatched_traceback = failure_dispatcher.await_args.args + assert isinstance(dispatched_exception, RuntimeError) + assert isinstance(dispatched_traceback, str) + assert dispatched_logging_obj is dependencies.proxy_logging.pre_call_hook.await_args.kwargs["data"]["litellm_logging_obj"] + dependencies.success_handler.pass_through_async_success_handler.assert_not_called() + dependencies.proxy_logging.post_call_failure_hook.assert_not_awaited() + finally: + cleanup() + await fake_client.aclose() + + +@pytest.mark.asyncio +async def test_websocket_passthrough_dispatches_team_failure_callback_for_pre_call_error(): + failure_dispatcher = AsyncMock() websocket = _client_websocket(AsyncMock(return_value={"type": "websocket.disconnect"})) + dependencies = _WebSocketPassthroughDependencies(RecordingUpstreamWebSocket()) + dependencies.proxy_logging.pre_call_hook = AsyncMock(side_effect=RuntimeError("pre-call blocked")) - with ExitStack() as stack: - stack.enter_context(patch("litellm.proxy.proxy_server.proxy_config", MagicMock())) - stack.enter_context( - patch.object( - litellm, - "_known_custom_logger_compatible_callbacks", - ["langfuse_otel"], - ) - ) - stack.enter_context( - patch( - "litellm.litellm_core_utils.litellm_logging._init_custom_logger_compatible_class", - return_value=failure_callback_recorder, - ) - ) - with _patched_websocket_passthrough_environment(ClosingUpstreamWebSocket(RuntimeError("upstream failed"))): - await websocket_passthrough_request( - websocket=websocket, - target="wss://upstream.test/v1/realtime", - custom_headers={}, - user_api_key_dict=_team_scoped_langfuse_otel_key(callback_type="failure"), - endpoint="/openai/v1/realtime", - ) + await websocket_passthrough_request( + websocket=websocket, + target="wss://upstream.test/v1/realtime", + custom_headers={}, + user_api_key_dict=_team_scoped_langfuse_otel_key(callback_type="failure"), + endpoint="/openai/v1/realtime", + proxy_config=MagicMock(), + dynamic_failure_dispatcher=failure_dispatcher, + **dependencies.kwargs(), + ) - assert len(failure_callback_recorder.failure_event_kwargs) == 1 - assert failure_callback_recorder.failure_event_kwargs[0]["additional_args"]["headers"] == {} + dispatched_logging_obj, dispatched_exception, dispatched_traceback = failure_dispatcher.await_args.args + assert isinstance(dispatched_exception, RuntimeError) + assert isinstance(dispatched_traceback, str) + assert dispatched_logging_obj.model_call_details["additional_args"]["headers"] == {} + dependencies.proxy_logging.pre_call_hook.assert_awaited_once_with( + user_api_key_dict=_team_scoped_langfuse_otel_key(callback_type="failure"), + data={}, + call_type="pass_through_endpoint", + ) + dependencies.proxy_logging.post_call_failure_hook.assert_awaited_once() + websocket.close.assert_awaited_once_with(code=1011, reason="WebSocket passthrough error") @pytest.mark.asyncio @@ -5629,62 +5747,39 @@ async def test_websocket_passthrough_dispatches_team_failure_callback_when_upstr from websockets.exceptions import InvalidStatus from websockets.http11 import Response as WebSocketResponse - failure_callback_recorder = _FailureCallbackRecorder() + failure_dispatcher = AsyncMock() websocket = _client_websocket(AsyncMock(return_value={"type": "websocket.disconnect"})) upstream_response = WebSocketResponse( status_code=403, reason_phrase="Forbidden", headers=WebSocketHeaders(), ) + dependencies = _WebSocketPassthroughDependencies(RecordingUpstreamWebSocket()) + dependencies.connect_factory.side_effect = InvalidStatus(upstream_response) - with ExitStack() as stack: - mock_proxy_logging = stack.enter_context(patch("litellm.proxy.proxy_server.proxy_logging_obj")) - mock_proxy_logging.pre_call_hook = AsyncMock(return_value={}) - mock_proxy_logging.post_call_failure_hook = AsyncMock() - stack.enter_context( - patch( - "litellm.proxy.pass_through_endpoints.pass_through_endpoints.connect", - side_effect=InvalidStatus(upstream_response), - ) - ) - mock_worker = stack.enter_context( - patch("litellm.proxy.pass_through_endpoints.pass_through_endpoints.GLOBAL_LOGGING_WORKER") - ) - mock_worker.ensure_initialized_and_enqueue = MagicMock( - side_effect=lambda async_coroutine: async_coroutine.close() - ) - stack.enter_context(patch("litellm.proxy.proxy_server.proxy_config", MagicMock())) - stack.enter_context( - patch.object( - litellm, - "_known_custom_logger_compatible_callbacks", - ["langfuse_otel"], - ) - ) - stack.enter_context( - patch( - "litellm.litellm_core_utils.litellm_logging._init_custom_logger_compatible_class", - return_value=failure_callback_recorder, - ) - ) + await websocket_passthrough_request( + websocket=websocket, + target="wss://upstream.test/v1/realtime", + custom_headers={}, + user_api_key_dict=_team_scoped_langfuse_otel_key(callback_type="failure"), + endpoint="/openai/v1/realtime", + proxy_config=MagicMock(), + dynamic_failure_dispatcher=failure_dispatcher, + **dependencies.kwargs(), + ) - await websocket_passthrough_request( - websocket=websocket, - target="wss://upstream.test/v1/realtime", - custom_headers={}, - user_api_key_dict=_team_scoped_langfuse_otel_key(callback_type="failure"), - endpoint="/openai/v1/realtime", - ) - - assert len(failure_callback_recorder.failure_event_kwargs) == 1 - mock_proxy_logging.post_call_failure_hook.assert_awaited_once() + dispatched_logging_obj, dispatched_exception, dispatched_traceback = failure_dispatcher.await_args.args + assert isinstance(dispatched_exception, InvalidStatus) + assert isinstance(dispatched_traceback, str) + assert dispatched_logging_obj.model_call_details["additional_args"]["headers"] == {} + dependencies.proxy_logging.post_call_failure_hook.assert_awaited_once() websocket.close.assert_awaited_once_with(code=1011, reason="Upstream connection rejected") @pytest.mark.asyncio async def test_websocket_passthrough_forwards_credentials_without_exposing_them_to_failure_callback(): - failure_callback_recorder = _FailureCallbackRecorder() - websocket = _client_websocket(AsyncMock(return_value={"type": "websocket.disconnect"})) + failure_dispatcher = AsyncMock() + websocket = _client_websocket(_pending_receive) websocket.headers = Headers( { "Authorization": "Bearer provider-secret", @@ -5693,64 +5788,31 @@ async def test_websocket_passthrough_forwards_credentials_without_exposing_them_ } ) upstream_ws = ClosingUpstreamWebSocket(RuntimeError("upstream failed")) + dependencies = _WebSocketPassthroughDependencies(upstream_ws) - with ExitStack() as stack: - mock_proxy_logging = stack.enter_context(patch("litellm.proxy.proxy_server.proxy_logging_obj")) - mock_proxy_logging.pre_call_hook = AsyncMock(return_value={}) - mock_proxy_logging.post_call_failure_hook = AsyncMock() - mock_connect = stack.enter_context( - patch( - "litellm.proxy.pass_through_endpoints.pass_through_endpoints.connect", - return_value=FakeUpstreamConnect(upstream_ws), - ) - ) - mock_worker = stack.enter_context( - patch("litellm.proxy.pass_through_endpoints.pass_through_endpoints.GLOBAL_LOGGING_WORKER") - ) - mock_worker.ensure_initialized_and_enqueue = MagicMock( - side_effect=lambda async_coroutine: async_coroutine.close() - ) - stack.enter_context(patch("litellm.proxy.proxy_server.proxy_config", MagicMock())) - stack.enter_context( - patch("litellm.proxy.proxy_server.general_settings", {"litellm_key_header_name": "x-custom-litellm-key"}) - ) - stack.enter_context( - patch.object( - litellm, - "_known_custom_logger_compatible_callbacks", - ["langfuse_otel"], - ) - ) - stack.enter_context( - patch( - "litellm.litellm_core_utils.litellm_logging._init_custom_logger_compatible_class", - return_value=failure_callback_recorder, - ) - ) + await websocket_passthrough_request( + websocket=websocket, + target="wss://upstream.test/v1/realtime", + custom_headers={}, + user_api_key_dict=_team_scoped_langfuse_otel_key(callback_type="failure"), + endpoint="/openai/v1/realtime", + forward_headers=True, + proxy_config=MagicMock(), + dynamic_failure_dispatcher=failure_dispatcher, + **dependencies.kwargs(), + ) - await websocket_passthrough_request( - websocket=websocket, - target="wss://upstream.test/v1/realtime", - custom_headers={"X-Custom-LiteLLM-Key": "virtual-key-secret"}, - user_api_key_dict=_team_scoped_langfuse_otel_key(callback_type="failure"), - endpoint="/openai/v1/realtime", - forward_headers=True, - ) - - assert len(failure_callback_recorder.failure_event_kwargs) == 1 - callback_kwargs = failure_callback_recorder.failure_event_kwargs[0] - callback_headers = callback_kwargs["additional_args"]["headers"] - assert callback_headers == { - "X-Custom-LiteLLM-Key": "***REDACTED***", + dispatched_logging_obj, dispatched_exception, dispatched_traceback = failure_dispatcher.await_args.args + assert isinstance(dispatched_exception, RuntimeError) + assert isinstance(dispatched_traceback, str) + assert dispatched_logging_obj.model_call_details["additional_args"]["headers"] == { "authorization": "***REDACTED***", "x-api-key": "***REDACTED***", } - serialized_callback_kwargs = json.dumps(callback_kwargs, default=str) - assert "provider-secret" not in serialized_callback_kwargs - assert "provider-api-key" not in serialized_callback_kwargs - assert "virtual-key-secret" not in serialized_callback_kwargs - assert mock_connect.call_args.kwargs["additional_headers"] == { - "X-Custom-LiteLLM-Key": "virtual-key-secret", + serialized_logging_details = json.dumps(dispatched_logging_obj.model_call_details, default=str) + assert "provider-secret" not in serialized_logging_details + assert "provider-api-key" not in serialized_logging_details + assert dependencies.connect_factory.call_args.kwargs["additional_headers"] == { "authorization": "Bearer provider-secret", "x-api-key": "provider-api-key", } diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_streaming_handler_interrupt.py b/tests/test_litellm/proxy/pass_through_endpoints/test_streaming_handler_interrupt.py index 1d82a5dfc6e..a5c9591900c 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_streaming_handler_interrupt.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_streaming_handler_interrupt.py @@ -9,55 +9,59 @@ import httpx import pytest import litellm -from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER from litellm.proxy.pass_through_endpoints.streaming_handler import ( PassThroughStreamingHandler, ) from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType -def _make_streaming_response(chunks): +def _make_streaming_response(chunks, error: Exception | None = None): mock = MagicMock(spec=httpx.Response) mock.status_code = 200 async def _aiter_bytes(): for c in chunks: yield c + if error is not None: + raise error mock.aiter_bytes = _aiter_bytes return mock +class _RecordingLoggingWorker: + def __init__(self): + self.coroutines = [] + + def ensure_initialized_and_enqueue(self, async_coroutine): + self.coroutines.append(async_coroutine) + async_coroutine.close() + + @pytest.mark.asyncio async def test_chunk_processor_logs_on_normal_completion(): chunks = [b"chunk-1", b"chunk-2", b"chunk-3"] response = _make_streaming_response(chunks) + route_logging = AsyncMock() + logging_worker = _RecordingLoggingWorker() - mock_logging_obj = MagicMock() - mock_passthrough_handler = MagicMock() - - with patch.object( - PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", - new=AsyncMock(), - ) as mock_route: - received = [] - async for chunk in PassThroughStreamingHandler.chunk_processor( - response=response, - request_body={"model": "claude-3-haiku"}, - litellm_logging_obj=mock_logging_obj, - endpoint_type=EndpointType.GENERIC, - start_time=datetime.now(), - passthrough_success_handler_obj=mock_passthrough_handler, - url_route="/bedrock/model/claude/invoke-with-response-stream", - ): - received.append(chunk) - - await asyncio.sleep(0) + received = [] + async for chunk in PassThroughStreamingHandler.chunk_processor( + response=response, + request_body={"model": "claude-3-haiku"}, + litellm_logging_obj=MagicMock(), + endpoint_type=EndpointType.GENERIC, + start_time=datetime.now(), + passthrough_success_handler_obj=MagicMock(), + url_route="/bedrock/model/claude/invoke-with-response-stream", + route_streaming_logging=route_logging, + logging_worker=logging_worker, + ): + received.append(chunk) assert received == chunks - mock_route.assert_called_once() - call_kwargs = mock_route.call_args.kwargs + assert len(logging_worker.coroutines) == 1 + call_kwargs = route_logging.call_args.kwargs assert call_kwargs["raw_bytes"] == chunks @@ -65,36 +69,84 @@ async def test_chunk_processor_logs_on_normal_completion(): async def test_chunk_processor_logs_on_client_disconnect(): chunks = [b"event-1", b"event-2", b"event-3"] response = _make_streaming_response(chunks) + route_logging = AsyncMock() + logging_worker = _RecordingLoggingWorker() - mock_logging_obj = MagicMock() - mock_passthrough_handler = MagicMock() + gen = PassThroughStreamingHandler.chunk_processor( + response=response, + request_body={"model": "claude-3-haiku"}, + litellm_logging_obj=MagicMock(), + endpoint_type=EndpointType.GENERIC, + start_time=datetime.now(), + passthrough_success_handler_obj=MagicMock(), + url_route="/bedrock/model/claude/invoke-with-response-stream", + route_streaming_logging=route_logging, + logging_worker=logging_worker, + ) - with patch.object( - PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", - new=AsyncMock(), - ) as mock_route: - gen = PassThroughStreamingHandler.chunk_processor( - response=response, - request_body={"model": "claude-3-haiku"}, - litellm_logging_obj=mock_logging_obj, - endpoint_type=EndpointType.GENERIC, - start_time=datetime.now(), - passthrough_success_handler_obj=mock_passthrough_handler, - url_route="/bedrock/model/claude/invoke-with-response-stream", - ) - - first = await gen.__anext__() - await gen.aclose() - - await asyncio.sleep(0) + first = await gen.__anext__() + await gen.aclose() assert first == chunks[0] - mock_route.assert_called_once() - call_kwargs = mock_route.call_args.kwargs + assert len(logging_worker.coroutines) == 1 + call_kwargs = route_logging.call_args.kwargs assert call_kwargs["raw_bytes"] == [chunks[0]] +@pytest.mark.asyncio +async def test_chunk_processor_dispatches_stream_read_error_without_success_logging(): + error = RuntimeError("upstream stream failed") + response = _make_streaming_response([b"partial"], error=error) + dispatcher = AsyncMock() + + route_logging = AsyncMock() + logging_obj = MagicMock() + logging_obj.model_call_details = {} + generator = PassThroughStreamingHandler.chunk_processor( + response=response, + request_body={"model": "claude-3-haiku"}, + litellm_logging_obj=logging_obj, + endpoint_type=EndpointType.GENERIC, + start_time=datetime.now(), + passthrough_success_handler_obj=MagicMock(), + url_route="/bedrock/model/claude/invoke-with-response-stream", + route_streaming_logging=route_logging, + dynamic_failure_dispatcher=dispatcher, + logging_worker=_RecordingLoggingWorker(), + ) + + assert await generator.__anext__() == b"partial" + with pytest.raises(RuntimeError, match="upstream stream failed"): + await generator.__anext__() + + dispatcher.assert_awaited_once_with(logging_obj, error) + route_logging.assert_not_called() + + +@pytest.mark.asyncio +async def test_chunk_processor_does_not_schedule_success_logging_for_upstream_error_without_dispatcher(): + response = _make_streaming_response([b"partial"], error=RuntimeError("upstream stream failed")) + + route_logging = AsyncMock() + generator = PassThroughStreamingHandler.chunk_processor( + response=response, + request_body={"model": "claude-3-haiku"}, + litellm_logging_obj=MagicMock(), + endpoint_type=EndpointType.GENERIC, + start_time=datetime.now(), + passthrough_success_handler_obj=MagicMock(), + url_route="/bedrock/model/claude/invoke-with-response-stream", + route_streaming_logging=route_logging, + logging_worker=_RecordingLoggingWorker(), + ) + + assert await generator.__anext__() == b"partial" + with pytest.raises(RuntimeError, match="upstream stream failed"): + await generator.__anext__() + + route_logging.assert_not_called() + + @pytest.mark.asyncio async def test_chunk_processor_does_not_schedule_success_logging_for_upstream_error(): """A 4xx/5xx upstream response is already logged as a failure by the caller @@ -104,58 +156,48 @@ async def test_chunk_processor_does_not_schedule_success_logging_for_upstream_er response = _make_streaming_response(chunks) response.status_code = 403 - mock_logging_obj = MagicMock() - mock_passthrough_handler = MagicMock() - - with patch.object( - PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", - new=AsyncMock(), - ) as mock_route: - received = [] - async for chunk in PassThroughStreamingHandler.chunk_processor( - response=response, - request_body={"model": "claude-3-haiku"}, - litellm_logging_obj=mock_logging_obj, - endpoint_type=EndpointType.GENERIC, - start_time=datetime.now(), - passthrough_success_handler_obj=mock_passthrough_handler, - url_route="/bedrock/model/claude/invoke-with-response-stream", - ): - received.append(chunk) - - await asyncio.sleep(0) + route_logging = AsyncMock() + logging_worker = _RecordingLoggingWorker() + received = [] + async for chunk in PassThroughStreamingHandler.chunk_processor( + response=response, + request_body={"model": "claude-3-haiku"}, + litellm_logging_obj=MagicMock(), + endpoint_type=EndpointType.GENERIC, + start_time=datetime.now(), + passthrough_success_handler_obj=MagicMock(), + url_route="/bedrock/model/claude/invoke-with-response-stream", + route_streaming_logging=route_logging, + logging_worker=logging_worker, + ): + received.append(chunk) assert received == chunks - mock_route.assert_not_called() + route_logging.assert_not_called() @pytest.mark.asyncio async def test_chunk_processor_does_not_schedule_logging_when_no_chunks(): response = _make_streaming_response([]) - mock_logging_obj = MagicMock() - mock_passthrough_handler = MagicMock() - - with patch.object( - PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", - new=AsyncMock(), - ) as mock_route: - received = [] - async for chunk in PassThroughStreamingHandler.chunk_processor( - response=response, - request_body={"model": "claude-3-haiku"}, - litellm_logging_obj=mock_logging_obj, - endpoint_type=EndpointType.GENERIC, - start_time=datetime.now(), - passthrough_success_handler_obj=mock_passthrough_handler, - url_route="/bedrock/model/claude/invoke-with-response-stream", - ): - received.append(chunk) + route_logging = AsyncMock() + logging_worker = _RecordingLoggingWorker() + received = [] + async for chunk in PassThroughStreamingHandler.chunk_processor( + response=response, + request_body={"model": "claude-3-haiku"}, + litellm_logging_obj=MagicMock(), + endpoint_type=EndpointType.GENERIC, + start_time=datetime.now(), + passthrough_success_handler_obj=MagicMock(), + url_route="/bedrock/model/claude/invoke-with-response-stream", + route_streaming_logging=route_logging, + logging_worker=logging_worker, + ): + received.append(chunk) assert received == [] - mock_route.assert_not_called() + route_logging.assert_not_called() @pytest.mark.asyncio @@ -167,39 +209,24 @@ async def test_chunk_processor_routes_logging_through_logging_worker(): chunks = [b"chunk-1", b"chunk-2"] response = _make_streaming_response(chunks) - enqueued = [] - - def _capture(async_coroutine): - enqueued.append(async_coroutine) - async_coroutine.close() - - with ( - patch.object( - PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", - new=AsyncMock(), - ), - patch.object( - GLOBAL_LOGGING_WORKER, - "ensure_initialized_and_enqueue", - side_effect=_capture, - ) as mock_enqueue, + logging_worker = _RecordingLoggingWorker() + received = [] + async for chunk in PassThroughStreamingHandler.chunk_processor( + response=response, + request_body={"model": "claude-3-haiku"}, + litellm_logging_obj=MagicMock(), + endpoint_type=EndpointType.GENERIC, + start_time=datetime.now(), + passthrough_success_handler_obj=MagicMock(), + url_route="/bedrock/model/claude/invoke-with-response-stream", + route_streaming_logging=AsyncMock(), + logging_worker=logging_worker, ): - received = [] - async for chunk in PassThroughStreamingHandler.chunk_processor( - response=response, - request_body={"model": "claude-3-haiku"}, - litellm_logging_obj=MagicMock(), - endpoint_type=EndpointType.GENERIC, - start_time=datetime.now(), - passthrough_success_handler_obj=MagicMock(), - url_route="/bedrock/model/claude/invoke-with-response-stream", - ): - received.append(chunk) + received.append(chunk) assert received == chunks - mock_enqueue.assert_called_once() - assert asyncio.iscoroutine(enqueued[0]) + assert len(logging_worker.coroutines) == 1 + assert asyncio.iscoroutine(logging_worker.coroutines[0]) @pytest.mark.asyncio @@ -209,38 +236,23 @@ async def test_chunk_processor_routes_logging_through_logging_worker_on_disconne chunks = [b"event-1", b"event-2", b"event-3"] response = _make_streaming_response(chunks) - enqueued = [] + logging_worker = _RecordingLoggingWorker() + gen = PassThroughStreamingHandler.chunk_processor( + response=response, + request_body={"model": "claude-3-haiku"}, + litellm_logging_obj=MagicMock(), + endpoint_type=EndpointType.GENERIC, + start_time=datetime.now(), + passthrough_success_handler_obj=MagicMock(), + url_route="/bedrock/model/claude/invoke-with-response-stream", + route_streaming_logging=AsyncMock(), + logging_worker=logging_worker, + ) + await gen.__anext__() + await gen.aclose() - def _capture(async_coroutine): - enqueued.append(async_coroutine) - async_coroutine.close() - - with ( - patch.object( - PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", - new=AsyncMock(), - ), - patch.object( - GLOBAL_LOGGING_WORKER, - "ensure_initialized_and_enqueue", - side_effect=_capture, - ) as mock_enqueue, - ): - gen = PassThroughStreamingHandler.chunk_processor( - response=response, - request_body={"model": "claude-3-haiku"}, - litellm_logging_obj=MagicMock(), - endpoint_type=EndpointType.GENERIC, - start_time=datetime.now(), - passthrough_success_handler_obj=MagicMock(), - url_route="/bedrock/model/claude/invoke-with-response-stream", - ) - await gen.__anext__() - await gen.aclose() - - mock_enqueue.assert_called_once() - assert asyncio.iscoroutine(enqueued[0]) + assert len(logging_worker.coroutines) == 1 + assert asyncio.iscoroutine(logging_worker.coroutines[0]) def _logging_obj_with_write_once_cst(): @@ -266,26 +278,19 @@ async def test_chunk_processor_stamps_completion_start_time_on_first_chunk(): response = _make_streaming_response(chunks) mock_logging_obj = _logging_obj_with_write_once_cst() - mock_passthrough_handler = MagicMock() - - with patch.object( - PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", - new=AsyncMock(), + received = [] + async for chunk in PassThroughStreamingHandler.chunk_processor( + response=response, + request_body={"model": "claude-haiku-4-5"}, + litellm_logging_obj=mock_logging_obj, + endpoint_type=EndpointType.ANTHROPIC, + start_time=datetime.now(), + passthrough_success_handler_obj=MagicMock(), + url_route="/v1/messages", + route_streaming_logging=AsyncMock(), + logging_worker=_RecordingLoggingWorker(), ): - received = [] - async for chunk in PassThroughStreamingHandler.chunk_processor( - response=response, - request_body={"model": "claude-haiku-4-5"}, - litellm_logging_obj=mock_logging_obj, - endpoint_type=EndpointType.ANTHROPIC, - start_time=datetime.now(), - passthrough_success_handler_obj=mock_passthrough_handler, - url_route="/v1/messages", - ): - received.append(chunk) - - await asyncio.sleep(0) + received.append(chunk) assert received == chunks mock_logging_obj._update_completion_start_time.assert_called_once() @@ -305,23 +310,19 @@ async def test_chunk_processor_does_not_reset_completion_start_time_on_later_chu # Simulate first-chunk stamp having already landed (e.g. under contention or a # prior wrapper that already set it): later chunks must be no-ops. mock_logging_obj.completion_start_time = real_first - mock_passthrough_handler = MagicMock() - with patch.object( - PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", - new=AsyncMock(), + async for _ in PassThroughStreamingHandler.chunk_processor( + response=response, + request_body={"model": "claude-haiku-4-5"}, + litellm_logging_obj=mock_logging_obj, + endpoint_type=EndpointType.ANTHROPIC, + start_time=datetime.now(), + passthrough_success_handler_obj=MagicMock(), + url_route="/v1/messages", + route_streaming_logging=AsyncMock(), + logging_worker=_RecordingLoggingWorker(), ): - async for _ in PassThroughStreamingHandler.chunk_processor( - response=response, - request_body={"model": "claude-haiku-4-5"}, - litellm_logging_obj=mock_logging_obj, - endpoint_type=EndpointType.ANTHROPIC, - start_time=datetime.now(), - passthrough_success_handler_obj=mock_passthrough_handler, - url_route="/v1/messages", - ): - pass + pass mock_logging_obj._update_completion_start_time.assert_not_called() assert mock_logging_obj.completion_start_time == real_first @@ -337,30 +338,30 @@ async def test_chunk_processor_stamps_completion_start_time_on_cost_injection_pa mock_logging_obj = _logging_obj_with_write_once_cst() mock_logging_obj.model_call_details = {"model": "claude-haiku-4-5"} - mock_passthrough_handler = MagicMock() + route_logging = AsyncMock() + logging_worker = _RecordingLoggingWorker() original = getattr(litellm_mod, "include_cost_in_streaming_usage", False) litellm_mod.include_cost_in_streaming_usage = True try: - with patch.object( - PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", - new=AsyncMock(), + async for _ in PassThroughStreamingHandler.chunk_processor( + response=response, + request_body={"model": "claude-haiku-4-5"}, + litellm_logging_obj=mock_logging_obj, + endpoint_type=EndpointType.ANTHROPIC, + start_time=datetime.now(), + passthrough_success_handler_obj=MagicMock(), + url_route="/v1/messages", + route_streaming_logging=route_logging, + logging_worker=logging_worker, ): - async for _ in PassThroughStreamingHandler.chunk_processor( - response=response, - request_body={"model": "claude-haiku-4-5"}, - litellm_logging_obj=mock_logging_obj, - endpoint_type=EndpointType.ANTHROPIC, - start_time=datetime.now(), - passthrough_success_handler_obj=mock_passthrough_handler, - url_route="/v1/messages", - ): - pass + pass finally: litellm_mod.include_cost_in_streaming_usage = original mock_logging_obj._update_completion_start_time.assert_called_once() + route_logging.assert_called_once() + assert len(logging_worker.coroutines) == 1 def _openai_passthrough_stream_chunks(): @@ -393,6 +394,7 @@ async def _collect_openai_passthrough_chunks(chunks, endpoint_type): passthrough_success_handler_obj=MagicMock(), url_route="/openai/v1/chat/completions", route_streaming_logging=AsyncMock(), + logging_worker=_RecordingLoggingWorker(), ): received.append(chunk) await asyncio.sleep(0) diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 959b2eada25..36d7e51fb6a 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -3,7 +3,7 @@ "limit": 22733 }, "LIT002": { - "limit": 26860 + "limit": 26859 }, "LIT003": { "limit": 269 @@ -27,7 +27,7 @@ "limit": 0 }, "LIT010": { - "limit": 16616 + "limit": 16612 }, "LIT011": { "limit": 5583