Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -247,6 +247,8 @@ def __init__(
self.audio_queue: asyncio.Queue[bytes | str] = asyncio.Queue()

self._reported_duration_ms = 0
self._num_output_tokens = 0
self._reported_output_tokens = 0

def _ensure_session(self) -> aiohttp.ClientSession:
"""Get or create an aiohttp ClientSession for WebSocket connections."""
Expand Down Expand Up @@ -306,24 +308,29 @@ async def _connect_ws(self) -> aiohttp.ClientWebSocketResponse:
logger.debug("Soniox Speech-to-Text API connection established!")

self._reported_duration_ms = 0
self._num_output_tokens = 0
self._reported_output_tokens = 0
self.audio_queue = asyncio.Queue()
return ws

def _report_processed_audio_duration(self, total_audio_proc_ms: float) -> None:
"""Report the total audio duration processed by the STT engine."""
"""Report the audio duration and output tokens processed by the STT engine."""
to_report_ms = total_audio_proc_ms - self._reported_duration_ms
if to_report_ms <= 0:
to_report_output_tokens = self._num_output_tokens - self._reported_output_tokens
if to_report_ms <= 0 and to_report_output_tokens <= 0:
return

usage_event = stt.SpeechEvent(
type=stt.SpeechEventType.RECOGNITION_USAGE,
alternatives=[],
recognition_usage=stt.RecognitionUsage(
audio_duration=to_report_ms / 1000,
audio_duration=max(to_report_ms, 0) / 1000,
output_tokens=max(to_report_output_tokens, 0),
),
)
self._event_ch.send_nowait(usage_event)
self._reported_duration_ms = int(total_audio_proc_ms)
self._reported_duration_ms = max(int(total_audio_proc_ms), self._reported_duration_ms)
self._reported_output_tokens = self._num_output_tokens

async def _run(self) -> None:
"""Manage connection lifecycle, spawning tasks and handling reconnection."""
Expand Down Expand Up @@ -523,6 +530,11 @@ def send_endpoint_transcript() -> None:
# 1) process tokens: accumulate final/non-final,
# flush immediately on endpoint tokens.
for token in tokens:
# Non-final tokens are re-sent as final later, so only
# final tokens count toward output usage. <end>/<fin>
# markers are control tokens, not transcription output.
if token["is_final"] and not is_end_token(token):
self._num_output_tokens += 1
is_translated = token.get("translation_status") == "translation"
if is_translation_mode and not is_end_token(token) and not is_translated:
# Original-language token: capture text for source_text only.
Expand Down
120 changes: 120 additions & 0 deletions tests/test_plugin_soniox_stt.py
Original file line number Diff line number Diff line change
Expand Up @@ -479,6 +479,126 @@ async def test_interim_transcript_no_translation_populates_source_runs():
assert sd.target_texts is None


# --- RECOGNITION_USAGE: output tokens ---------------------------------------


async def test_usage_event_reports_final_tokens_as_output_tokens():
"""Each final transcription token counts once toward `output_tokens`;
the `<end>` control token does not."""
stream = _make_stream(translation=None)

messages = [
{
"tokens": [
_final_token("Hello", "en"),
_final_token(" world.", "en"),
END_TOKEN_FINAL,
],
"total_audio_proc_ms": 500,
}
]

events = await _drive_recv(stream, messages, expect_events=3)
usage = next(e for e in events if e.type == SpeechEventType.RECOGNITION_USAGE)

assert usage.recognition_usage is not None
assert usage.recognition_usage.audio_duration == 0.5
assert usage.recognition_usage.output_tokens == 2


async def test_usage_event_excludes_non_final_tokens():
"""Non-final tokens are re-sent as final later — counting them would
double-bill. Only the final occurrences count."""
stream = _make_stream(translation=None)

messages = [
{
"tokens": [
_nonfinal_token("Hel", "en"),
_nonfinal_token("lo", "en"),
],
"total_audio_proc_ms": 300,
},
{
"tokens": [
_final_token("Hello", "en"),
_final_token(" world.", "en"),
END_TOKEN_FINAL,
],
"total_audio_proc_ms": 800,
},
]

# message 1: START_OF_SPEECH + INTERIM; message 2: FINAL + END_OF_SPEECH + USAGE.
events = await _drive_recv(stream, messages, expect_events=5)
usage = next(e for e in events if e.type == SpeechEventType.RECOGNITION_USAGE)

assert usage.recognition_usage is not None
assert usage.recognition_usage.output_tokens == 2


async def test_usage_events_report_incremental_output_tokens():
"""A second endpoint reports only the tokens accumulated since the
previous usage event, mirroring the incremental audio_duration."""
stream = _make_stream(translation=None)

messages = [
{
"tokens": [
_final_token("First.", "en"),
END_TOKEN_FINAL,
],
"total_audio_proc_ms": 1000,
},
{
"tokens": [
_final_token("Second", "en"),
_final_token(" utterance.", "en"),
END_TOKEN_FINAL,
],
"total_audio_proc_ms": 2500,
},
]

# Per endpoint: START_OF_SPEECH is not emitted before endpoint flush
# (tokens are final-only), so each message yields FINAL + END_OF_SPEECH + USAGE.
events = await _drive_recv(stream, messages, expect_events=6)
usages = [e for e in events if e.type == SpeechEventType.RECOGNITION_USAGE]

assert len(usages) == 2
assert usages[0].recognition_usage is not None
assert usages[0].recognition_usage.audio_duration == 1.0
assert usages[0].recognition_usage.output_tokens == 1
assert usages[1].recognition_usage is not None
assert usages[1].recognition_usage.audio_duration == 1.5
assert usages[1].recognition_usage.output_tokens == 2


async def test_usage_event_counts_translation_tokens():
"""In translation mode both original and translated final tokens are
provider output and count toward `output_tokens`."""
from livekit.plugins.soniox.stt import TranslationConfig

stream = _make_stream(translation=TranslationConfig(type="one_way", target_language="es"))

messages = [
{
"tokens": [
_final_token("Hello world.", "en", translation_status="original"),
_final_token("Hola mundo.", "es", translation_status="translation"),
END_TOKEN_FINAL,
],
"total_audio_proc_ms": 500,
}
]

events = await _drive_recv(stream, messages, expect_events=3)
usage = next(e for e in events if e.type == SpeechEventType.RECOGNITION_USAGE)

assert usage.recognition_usage is not None
assert usage.recognition_usage.output_tokens == 2


async def test_recv_messages_raises_on_server_error_frame():
stream = _make_stream(translation=None)
stream._ws = _FakeWebSocket(
Expand Down