mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(proxy): harden passthrough callback dispatch
This commit is contained in:
parent
c8ae1cf4ed
commit
05093cb85e
8 changed files with 849 additions and 680 deletions
|
|
@ -84,7 +84,7 @@
|
|||
"limit": 56
|
||||
},
|
||||
"reportPrivateUsage": {
|
||||
"limit": 1806
|
||||
"limit": 1803
|
||||
},
|
||||
"reportRedeclaration": {
|
||||
"limit": 8
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -57,7 +57,7 @@
|
|||
"limit": 3
|
||||
},
|
||||
"BLE001": {
|
||||
"limit": 2917
|
||||
"limit": 2914
|
||||
},
|
||||
"C401": {
|
||||
"limit": 8
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue