fix(elevenlabs): fix import order, preserve char count across task cancellation

Sort lazy imports inside the endpoint to satisfy ruff isort (CI used
ruff 0.15.3 which enforces isort within function bodies).

Refactor _relay_client_to_upstream to accept an external char_totals
list that is appended to before each upstream.send(). This ensures that
character counts accumulated before a task cancellation are preserved:
when ElevenLabs sends isFinal and the upstream-to-client relay finishes
first, the client-to-upstream relay is cancelled but char_totals already
reflects every chunk forwarded up to that point, so the session is never
recorded at zero cost after real usage.
This commit is contained in:
javimp2003uma 2026-08-17 00:43:46 +02:00
parent 87bd639bc0
commit 357cfb3a79
3 changed files with 31 additions and 18 deletions

View file

@ -65,16 +65,21 @@ def build_elevenlabs_ws_url(
async def _relay_client_to_upstream(
client_ws: WebSocket,
upstream: ClientConnection,
) -> tuple[int, ...]:
chunk_lengths: list[int] = [] # mutable-ok: local accumulator, converted to immutable tuple on return
char_totals: list[int],
) -> None:
"""Forward client text chunks to the upstream ElevenLabs connection.
Characters are appended to `char_totals` *before* the upstream send so that
partial counts survive task cancellation (e.g. when ElevenLabs sends isFinal
before the client sends EOS).
"""
async for raw in client_ws.iter_text():
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", "") # rebind-ok: loop-body, rebound each iteration
chunk_lengths.append(len(text))
char_totals.append(len(text))
await upstream.send(json.dumps(dict(msg)))
if not text:
break
return tuple(chunk_lengths)
async def _relay_upstream_to_client(
@ -109,6 +114,8 @@ async def stream_input_tts(
back to the client as-is.
Returns the total number of text characters sent (used for per-character cost tracking).
Character counts are committed to an external accumulator before each upstream send,
so partial totals are preserved even if the relay is cancelled early.
"""
import websockets
@ -124,20 +131,24 @@ async def stream_input_tts(
if generation_config is not None:
bos["generation_config"] = generation_config
char_totals: list[int] = [] # mutable-ok: accumulator written before each send; survives task cancellation
async with websockets.connect(url, additional_headers={"xi-api-key": key}) as upstream:
await upstream.send(json.dumps(bos))
task_c2u: Final = asyncio.create_task(_relay_client_to_upstream(client_ws, upstream))
task_c2u: Final = asyncio.create_task(
_relay_client_to_upstream(client_ws, upstream, char_totals)
)
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.
# Prevents the client-to-upstream relay from hanging if ElevenLabs sends
# isFinal before the client sends EOS, or if either side disconnects.
# char_totals is written before each upstream.send(), so its contents
# reflect all text actually forwarded, even if task_c2u is cancelled mid-session.
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)
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)
return sum(char_totals)

View file

@ -10910,6 +10910,7 @@ async def elevenlabs_tts_stream_input_endpoint(
import httpx
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.llms.elevenlabs.text_to_speech.transformation import (
ElevenLabsTextToSpeechConfig,
)
@ -10917,7 +10918,6 @@ async def elevenlabs_tts_stream_input_endpoint(
VoiceSettings,
stream_input_tts,
)
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.types.llms.openai import HttpxBinaryResponseContent
elevenlabs_model: Final = model.removeprefix("elevenlabs/")

View file

@ -92,9 +92,10 @@ class TestRelayClientToUpstream:
upstream = AsyncMock()
upstream.send = AsyncMock()
result = await _relay_client_to_upstream(client_ws, upstream)
char_totals: list[int] = []
await _relay_client_to_upstream(client_ws, upstream, char_totals)
assert result == (7, 7, 0)
assert char_totals == [7, 7, 0]
assert upstream.send.call_count == 3
@pytest.mark.asyncio
@ -110,9 +111,10 @@ class TestRelayClientToUpstream:
upstream = AsyncMock()
upstream.send = AsyncMock()
result = await _relay_client_to_upstream(client_ws, upstream)
char_totals: list[int] = []
await _relay_client_to_upstream(client_ws, upstream, char_totals)
assert sum(result) == 6
assert sum(char_totals) == 6
assert upstream.send.call_count == 2
@pytest.mark.asyncio
@ -130,7 +132,7 @@ class TestRelayClientToUpstream:
upstream.send = capture_send
await _relay_client_to_upstream(client_ws, upstream)
await _relay_client_to_upstream(client_ws, upstream, [])
assert sent_payloads[0]["flush"] is True
assert sent_payloads[0]["try_trigger_generation"] is True