diff --git a/litellm/llms/elevenlabs/text_to_speech/ws_handler.py b/litellm/llms/elevenlabs/text_to_speech/ws_handler.py index 12ad5a0c1e8..12d3b8c6a29 100644 --- a/litellm/llms/elevenlabs/text_to_speech/ws_handler.py +++ b/litellm/llms/elevenlabs/text_to_speech/ws_handler.py @@ -3,6 +3,7 @@ from __future__ import annotations import asyncio import json from typing import TYPE_CHECKING, Final, Required +from urllib.parse import urlencode from pydantic import TypeAdapter from starlette.websockets import WebSocket @@ -57,7 +58,8 @@ def build_elevenlabs_ws_url( ws_base: Final = raw_base.replace("https://", "wss://").replace("http://", "ws://") encoded_voice: Final = encode_url_path_segment(voice_id, field_name="voice_id") path: Final = _WS_PATH.format(voice_id=encoded_voice) - return f"{ws_base}{path}?model_id={model}&output_format={output_format}" + query: Final = urlencode({"model_id": model, "output_format": output_format}) + return f"{ws_base}{path}?{query}" async def _relay_client_to_upstream( @@ -66,9 +68,9 @@ async def _relay_client_to_upstream( ) -> tuple[int, ...]: chunk_lengths: list[int] = [] # mutable-ok: local accumulator, converted to immutable tuple on return async for raw in client_ws.iter_text(): - msg = _CLIENT_MSG_ADAPTER.validate_json(raw) + msg = _CLIENT_MSG_ADAPTER.validate_json(raw) # rebind-ok: loop-body, rebound each iteration await upstream.send(json.dumps(dict(msg))) - text = msg.get("text", "") + text = msg.get("text", "") # rebind-ok: loop-body, rebound each iteration chunk_lengths.append(len(text)) if not text: break @@ -80,9 +82,9 @@ async def _relay_upstream_to_client( client_ws: WebSocket, ) -> None: async for raw in upstream: - payload = raw if isinstance(raw, str) else raw.decode() + payload = raw if isinstance(raw, str) else raw.decode() # rebind-ok: loop-body, rebound each iteration await client_ws.send_text(payload) - msg = _SERVER_MSG_ADAPTER.validate_json(payload) + msg = _SERVER_MSG_ADAPTER.validate_json(payload) # rebind-ok: loop-body, rebound each iteration if msg.get("isFinal"): break @@ -124,10 +126,18 @@ async def stream_input_tts( async with websockets.connect(url, additional_headers={"xi-api-key": key}) as upstream: await upstream.send(json.dumps(bos)) - results: Final = await asyncio.gather( - _relay_client_to_upstream(client_ws, upstream), - _relay_upstream_to_client(upstream, client_ws), - ) - chunk_lengths: Final = results[0] + task_c2u: Final = asyncio.create_task(_relay_client_to_upstream(client_ws, upstream)) + task_u2c: Final = asyncio.create_task(_relay_upstream_to_client(upstream, client_ws)) + + # Wait for whichever relay finishes first, then cancel the other. + # This prevents the client-to-upstream relay from hanging if ElevenLabs + # sends isFinal before the client sends EOS, or if either side disconnects. + await asyncio.wait({task_c2u, task_u2c}, return_when=asyncio.FIRST_COMPLETED) + task_c2u.cancel() + task_u2c.cancel() + results: Final = await asyncio.gather(task_c2u, task_u2c, return_exceptions=True) + + c2u_result: Final = results[0] + chunk_lengths: Final = c2u_result if isinstance(c2u_result, tuple) else () return sum(chunk_lengths) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index c8e1c19061a..79cfcd6a49e 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -10934,8 +10934,6 @@ async def elevenlabs_tts_stream_input_endpoint( await websocket.close(code=1008, reason=e.message[:120]) return - await websocket.accept() - elevenlabs_config: Final = ElevenLabsTextToSpeechConfig() voice_id: Final = elevenlabs_config._extract_voice_id(voice) @@ -10951,6 +10949,21 @@ async def elevenlabs_tts_stream_input_endpoint( } voice_settings: Final[VoiceSettings | None] = cast(VoiceSettings, raw_voice_settings) if raw_voice_settings else None # cast-ok: dict built from typed float query params; structural match is guaranteed + # Run guardrails and rate-limit checks before opening the upstream connection. + # This mirrors what the batch TTS endpoint does via proxy_logging_obj.pre_call_hook(). + _initial_hook_data: Final[dict[str, object]] = {"model": litellm_model, "user": user_api_key_dict.user_id} + try: + pre_call_data: Final = await proxy_logging_obj.pre_call_hook( + user_api_key_dict=user_api_key_dict, + data=_initial_hook_data, + call_type="aspeech", + ) + except Exception as pre_call_err: # noqa: BLE001 # guardrail block or rate-limit; close before accepting + await websocket.close(code=1008, reason=str(pre_call_err)[:120]) + return + + await websocket.accept() + start_time: Final = datetime.now(tz=timezone.utc) litellm_call_id: Final = str(uuid4()) @@ -10963,6 +10976,21 @@ async def elevenlabs_tts_stream_input_endpoint( litellm_call_id=litellm_call_id, function_id="elevenlabs_tts_stream_input", ) + # Attribute cost to the authenticated key/team so budget enforcement works. + logging_obj.update_environment_variables( + model=litellm_model, + user=user_api_key_dict.user_id, + optional_params={}, + litellm_params={ + "metadata": { + "user_api_key": user_api_key_dict.api_key, + "user_api_key_alias": user_api_key_dict.key_alias, + "user_api_key_user_id": user_api_key_dict.user_id, + "user_api_key_team_id": user_api_key_dict.team_id, + } + }, + custom_llm_provider="elevenlabs", + ) try: total_chars: Final = await stream_input_tts( @@ -10998,8 +11026,13 @@ async def elevenlabs_tts_stream_input_endpoint( end_time=end_time, cache_hit=False, ) - except Exception: # noqa: BLE001 # intentional: catch all session errors to ensure WS cleanup + except Exception as session_err: # noqa: BLE001 # intentional: catch all session errors to ensure WS cleanup verbose_proxy_logger.exception("ElevenLabs TTS stream-input error") + await proxy_logging_obj.post_call_failure_hook( + user_api_key_dict=user_api_key_dict, + original_exception=session_err, + request_data=pre_call_data, + ) try: await websocket.close(code=1011, reason="Internal server error") except Exception: # noqa: BLE001 # WS may already be closed; log and discard