mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
ee3ddfe91e
commit
d196796c37
2 changed files with 56 additions and 13 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue