Skip to content
Merged
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 @@ -76,6 +76,7 @@ class STTOptions:
keyterms: NotGivenOr[list[str]]
no_verbatim: bool
enable_logging: bool
previous_text: str | None


class STT(stt.STT):
Expand All @@ -96,6 +97,7 @@ def __init__(
keyterms: NotGivenOr[list[str]] = NOT_GIVEN,
no_verbatim: NotGivenOr[bool] = NOT_GIVEN,
enable_logging: bool = True,
previous_text: NotGivenOr[str] = NOT_GIVEN,
) -> None:
"""
Create a new instance of ElevenLabs STT.
Expand Down Expand Up @@ -123,6 +125,8 @@ def __init__(
Scribe v2 (batch) and Scribe v2 realtime. Default is False.
enable_logging (bool): Enable logging of the request. When set to false, zero retention
mode will be used. Defaults to True.
previous_text (NotGivenOr[str]): Preceding text context sent once on the first realtime
audio chunk to improve transcription accuracy. Only supported for Scribe v2 realtime.
"""

if is_given(model_id):
Expand Down Expand Up @@ -152,6 +156,13 @@ def __init__(
if not use_realtime and is_given(server_vad):
logger.warning("Server-side VAD is only supported for Scribe v2 realtime model")

resolved_previous_text = previous_text if is_given(previous_text) else None
if not use_realtime and resolved_previous_text is not None:
logger.warning(
"`previous_text` is only supported for Scribe v2 realtime model and will be ignored"
)
resolved_previous_text = None

super().__init__(
capabilities=STTCapabilities(
streaming=use_realtime,
Expand Down Expand Up @@ -179,6 +190,7 @@ def __init__(
keyterms=keyterms,
no_verbatim=no_verbatim if is_given(no_verbatim) else False,
enable_logging=enable_logging,
previous_text=resolved_previous_text,
)
self._session = http_session
self._streams = weakref.WeakSet[SpeechStream]()
Expand Down Expand Up @@ -490,6 +502,19 @@ async def recv_task(ws: aiohttp.ClientWebSocketResponse) -> None:
while True:
try:
ws = await self._connect_ws()
if self._opts.previous_text:
# Must be the first input_audio_chunk on the connection.
await ws.send_str(
json.dumps(
{
"message_type": "input_audio_chunk",
"audio_base_64": "",
"commit": False,
"sample_rate": self._opts.sample_rate,
"previous_text": self._opts.previous_text,
}
)
)
tasks = [
asyncio.create_task(send_task(ws)),
asyncio.create_task(recv_task(ws)),
Expand Down
17 changes: 17 additions & 0 deletions tests/test_plugin_elevenlabs_stt.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,7 @@ def _new_stream(*, server_vad=NOT_GIVEN) -> elevenlabs_stt.SpeechStream:
keyterms=NOT_GIVEN,
no_verbatim=False,
enable_logging=True,
previous_text=None,
)
stream._language = None
stream._event_ch = _EventSink()
Expand Down Expand Up @@ -136,6 +137,22 @@ def test_enable_logging_can_be_disabled() -> None:
assert _stt(enable_logging=False)._opts.enable_logging is False


def test_previous_text_is_kept_for_realtime_model() -> None:
assert _stt(previous_text="prior context")._opts.previous_text == "prior context"


def test_previous_text_is_ignored_for_non_realtime_model(caplog: pytest.LogCaptureFixture) -> None:
with caplog.at_level("WARNING"):
instance = elevenlabs_stt.STT(
api_key="test-key",
model="scribe_v2",
previous_text="prior context",
)

assert instance._opts.previous_text is None
assert any("previous_text" in record.message for record in caplog.records)


@pytest.mark.parametrize(("enable_logging", "expected"), [(True, "true"), (False, "false")])
async def test_connect_ws_includes_enable_logging(enable_logging: bool, expected: str) -> None:
# enable_logging is a WebSocket query param. Verify it is forwarded to the
Expand Down