fix(proxy): harden passthrough callback dispatch

This commit is contained in:
Yucheng Zhu 2026-08-29 12:16:25 -07:00
parent c8ae1cf4ed
commit 05093cb85e
8 changed files with 849 additions and 680 deletions

View file

@ -84,7 +84,7 @@
"limit": 56
},
"reportPrivateUsage": {
"limit": 1806
"limit": 1803
},
"reportRedeclaration": {
"limit": 8

View file

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

View file

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

View file

@ -57,7 +57,7 @@
"limit": 3
},
"BLE001": {
"limit": 2917
"limit": 2914
},
"C401": {
"limit": 8

View file

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

View file

@ -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",
}

View file

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

View file

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