mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
87bd639bc0
commit
357cfb3a79
3 changed files with 31 additions and 18 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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/")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue