diff --git a/litellm/llms/elevenlabs/text_to_speech/ws_handler.py b/litellm/llms/elevenlabs/text_to_speech/ws_handler.py index 12d3b8c6a29..55503008f28 100644 --- a/litellm/llms/elevenlabs/text_to_speech/ws_handler.py +++ b/litellm/llms/elevenlabs/text_to_speech/ws_handler.py @@ -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) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index fde4d4cda3a..d09bd06e35c 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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/") diff --git a/tests/test_litellm/llms/elevenlabs/test_elevenlabs_ws_tts_handler.py b/tests/test_litellm/llms/elevenlabs/test_elevenlabs_ws_tts_handler.py index 6f7621a4c9c..d06f36ce87b 100644 --- a/tests/test_litellm/llms/elevenlabs/test_elevenlabs_ws_tts_handler.py +++ b/tests/test_litellm/llms/elevenlabs/test_elevenlabs_ws_tts_handler.py @@ -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