fix(elevenlabs): address relay hang, spend attribution, guardrails and URL injection

Relay hang: replace asyncio.gather() with asyncio.wait(FIRST_COMPLETED)
plus explicit task cancellation so the client-to-upstream relay does not
block indefinitely when ElevenLabs sends isFinal before the client
closes with {"text": ""}.

Spend attribution: call logging_obj.update_environment_variables() with
the authenticated key's api_key, key_alias, user_id and team_id so
budget enforcement callbacks receive the correct key context.

Guardrails / rate limits: run proxy_logging_obj.pre_call_hook() before
accepting the WebSocket connection, matching what the batch TTS endpoint
does. A guardrail block or rate-limit hit closes the socket with 1008
before the ElevenLabs upstream is opened.

URL injection: use urllib.parse.urlencode() for model_id and
output_format query parameters so crafted values cannot inject extra
query string fields into the upstream URL.
This commit is contained in:
javimp2003uma 2026-08-16 22:25:39 +02:00
parent ee3ddfe91e
commit d196796c37
2 changed files with 56 additions and 13 deletions

View file

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

View file

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