Add LiteLLM speech-to-text support (#183073)

This commit is contained in:
Marcin Juraszek
2026-10-03 22:17:48 +02:00
committed by GitHub
parent 1c100dba4a
commit a88f859aae
7 changed files with 870 additions and 4 deletions
+1 -1
View File
@@ -5,7 +5,7 @@ from homeassistant.core import HomeAssistant
from .coordinator import LiteLLMConfigEntry, LiteLLMDataUpdateCoordinator
PLATFORMS = [Platform.CONVERSATION]
PLATFORMS = [Platform.CONVERSATION, Platform.STT]
async def async_setup_entry(hass: HomeAssistant, entry: LiteLLMConfigEntry) -> bool:
+120 -1
View File
@@ -31,6 +31,9 @@ from homeassistant.helpers.selector import (
from .const import (
CONF_PROMPT,
CONF_STT_CUSTOM_PROMPT_KEYWORDS,
CONF_STT_KEYWORDS,
CONF_STT_PROMPT,
DOMAIN,
PLACEHOLDER_API_KEY,
RECOMMENDED_CONVERSATION_OPTIONS,
@@ -90,7 +93,10 @@ class LiteLLMConfigFlow(ConfigFlow, domain=DOMAIN):
cls, config_entry: ConfigEntry
) -> dict[str, type[ConfigSubentryFlow]]:
"""Return subentries supported by this handler."""
return {"conversation": ConversationFlowHandler}
return {
"conversation": ConversationFlowHandler,
"stt": STTFlowHandler,
}
@override
async def async_step_user(
@@ -252,3 +258,116 @@ class ConversationFlowHandler(LiteLLMSubentryFlowHandler):
}
),
)
class STTFlowHandler(LiteLLMSubentryFlowHandler):
"""Handle STT subentry flow."""
def __init__(self) -> None:
"""Initialize the subentry flow."""
super().__init__()
self.options: dict[str, Any] = {}
self.last_rendered_custom_prompt_keywords = False
@property
def _is_new(self) -> bool:
"""Return if this is a new subentry."""
return self.source == SOURCE_USER
async def async_step_user(
self, user_input: dict[str, Any] | None = None
) -> SubentryFlowResult:
"""User flow to create an STT entity."""
self.options = {CONF_STT_CUSTOM_PROMPT_KEYWORDS: False}
return await self.async_step_init(user_input)
async def async_step_reconfigure(
self, user_input: dict[str, Any] | None = None
) -> SubentryFlowResult:
"""Handle reconfiguration of an STT entity."""
self.options = self._get_reconfigure_subentry().data.copy()
return await self.async_step_init(user_input)
async def async_step_init(
self, user_input: dict[str, Any] | None = None
) -> SubentryFlowResult:
"""Manage STT configuration."""
if self._get_entry().state is not ConfigEntryState.LOADED:
return self.async_abort(reason="entry_not_loaded")
if user_input is not None:
custom_prompt_keywords = user_input[CONF_STT_CUSTOM_PROMPT_KEYWORDS]
if custom_prompt_keywords == self.last_rendered_custom_prompt_keywords:
self.options = user_input.copy()
for field in (CONF_STT_PROMPT, CONF_STT_KEYWORDS):
if not custom_prompt_keywords or not self.options.get(field):
self.options.pop(field, None)
if self._is_new:
return self.async_create_entry(
title=self.options[CONF_MODEL], data=self.options
)
return self.async_update_and_abort(
self._get_entry(),
self._get_reconfigure_subentry(),
title=self.options[CONF_MODEL],
data=self.options,
)
self.options = user_input
self.last_rendered_custom_prompt_keywords = custom_prompt_keywords
else:
self.last_rendered_custom_prompt_keywords = bool(
self.options.get(CONF_STT_CUSTOM_PROMPT_KEYWORDS, False)
)
try:
await self._fetch_models()
except InvalidAuth:
return self.async_abort(reason="invalid_auth")
except CannotConnect:
return self.async_abort(reason="cannot_connect")
except Exception:
_LOGGER.exception("Unexpected exception")
return self.async_abort(reason="unknown")
schema: dict[Any, Any] = {
probatio.Required(
CONF_MODEL, default=self.options.get(CONF_MODEL)
): SelectSelector(
SelectSelectorConfig(
options=[
SelectOptionDict(value=model, label=model)
for model in self.models
],
mode=SelectSelectorMode.DROPDOWN,
sort=True,
)
),
probatio.Required(
CONF_STT_CUSTOM_PROMPT_KEYWORDS,
default=self.options.get(CONF_STT_CUSTOM_PROMPT_KEYWORDS, False),
): bool,
}
if self.options.get(CONF_STT_CUSTOM_PROMPT_KEYWORDS):
schema.update(
{
probatio.Optional(
CONF_STT_PROMPT,
description={
"suggested_value": self.options.get(CONF_STT_PROMPT, "")
},
): TemplateSelector(),
probatio.Optional(
CONF_STT_KEYWORDS,
description={
"suggested_value": self.options.get(CONF_STT_KEYWORDS, "")
},
): TemplateSelector(),
}
)
return self.async_show_form(
step_id="init",
data_schema=probatio.Schema(schema),
)
@@ -7,6 +7,9 @@ from homeassistant.helpers import llm
DOMAIN = "litellm"
LOGGER = logging.getLogger(__package__)
CONF_STT_CUSTOM_PROMPT_KEYWORDS = "custom_prompt_keywords"
CONF_STT_KEYWORDS = "keywords"
CONF_STT_PROMPT = "prompt"
# LiteLLM proxies may run without authentication. The OpenAI client requires a
# non-empty API key, so we send a placeholder when the user did not provide one.
@@ -49,6 +49,36 @@
"description": "Configure the conversation agent"
}
}
},
"stt": {
"abort": {
"cannot_connect": "[%key:common::config_flow::error::cannot_connect%]",
"entry_not_loaded": "The main integration entry is not loaded. Please ensure the integration is loaded before reconfiguring.",
"invalid_auth": "[%key:common::config_flow::error::invalid_auth%]",
"unknown": "[%key:common::config_flow::error::unknown%]"
},
"entry_type": "Speech-to-text",
"initiate_flow": {
"reconfigure": "Reconfigure speech-to-text service",
"user": "Add speech-to-text service"
},
"step": {
"init": {
"data": {
"custom_prompt_keywords": "Custom prompt/keywords",
"keywords": "Keywords",
"model": "[%key:common::generic::model%]",
"prompt": "Prompt"
},
"data_description": {
"custom_prompt_keywords": "Enable optional prompt and keyword hints. The selected LiteLLM backend must support each parameter; unsupported parameters can cause transcription errors.",
"keywords": "Comma-separated words or phrases to help the model recognize names and other terms. Supports Home Assistant templates. OpenAI models may support this together with a prompt; other models may not.",
"model": "The model to use to transcribe speech",
"prompt": "Optional context or instructions for the transcription. Supports Home Assistant templates. The prompt must match the language of the audio and be supported by the selected model."
},
"description": "Select the speech-to-text model. The transcription language is selected in the Assist pipeline. LiteLLM does not report language capabilities for individual backend models."
}
}
}
}
}
+232
View File
@@ -0,0 +1,232 @@
"""Speech-to-text support for LiteLLM."""
from collections.abc import AsyncIterable
import io
from typing import override
import wave
from openai import (
APIConnectionError,
AuthenticationError,
Omit,
OpenAIError,
PermissionDeniedError,
omit,
)
from homeassistant.components import stt
from homeassistant.core import HomeAssistant
from homeassistant.exceptions import TemplateError
from homeassistant.helpers.entity_platform import AddConfigEntryEntitiesCallback
from homeassistant.helpers.template import Template
from .const import (
CONF_STT_CUSTOM_PROMPT_KEYWORDS,
CONF_STT_KEYWORDS,
CONF_STT_PROMPT,
LOGGER,
)
from .coordinator import LiteLLMConfigEntry
from .entity import LiteLLMEntity
PARALLEL_UPDATES = 0
async def async_setup_entry(
hass: HomeAssistant,
config_entry: LiteLLMConfigEntry,
async_add_entities: AddConfigEntryEntitiesCallback,
) -> None:
"""Set up LiteLLM STT entities."""
for subentry in config_entry.get_subentries_of_type("stt"):
async_add_entities(
[LiteLLMSTTEntity(config_entry, subentry)],
config_subentry_id=subentry.subentry_id,
)
class LiteLLMSTTEntity(LiteLLMEntity, stt.SpeechToTextEntity):
"""LiteLLM speech-to-text entity."""
@property
@override
def supported_languages(self) -> list[str]:
"""Return supported languages.
LiteLLM does not expose model-specific language capabilities; the
selected backend may reject the advertised language.
"""
return [
"af-ZA", # Afrikaans
"ar-SA", # Arabic
"hy-AM", # Armenian
"az-AZ", # Azerbaijani
"be-BY", # Belarusian
"bs-BA", # Bosnian
"bg-BG", # Bulgarian
"ca-ES", # Catalan
"zh-CN", # Chinese (Mandarin)
"hr-HR", # Croatian
"cs-CZ", # Czech
"da-DK", # Danish
"nl-NL", # Dutch
"en-US", # English
"et-EE", # Estonian
"fi-FI", # Finnish
"fr-FR", # French
"gl-ES", # Galician
"de-DE", # German
"el-GR", # Greek
"he-IL", # Hebrew
"hi-IN", # Hindi
"hu-HU", # Hungarian
"is-IS", # Icelandic
"id-ID", # Indonesian
"it-IT", # Italian
"ja-JP", # Japanese
"kn-IN", # Kannada
"kk-KZ", # Kazakh
"ko-KR", # Korean
"lv-LV", # Latvian
"lt-LT", # Lithuanian
"mk-MK", # Macedonian
"ms-MY", # Malay
"mr-IN", # Marathi
"mi-NZ", # Maori
"ne-NP", # Nepali
"no-NO", # Norwegian
"fa-IR", # Persian
"pl-PL", # Polish
"pt-PT", # Portuguese
"ro-RO", # Romanian
"ru-RU", # Russian
"sr-RS", # Serbian
"sk-SK", # Slovak
"sl-SI", # Slovenian
"es-ES", # Spanish
"sw-KE", # Swahili
"sv-SE", # Swedish
"fil-PH", # Tagalog (Filipino)
"ta-IN", # Tamil
"th-TH", # Thai
"tr-TR", # Turkish
"uk-UA", # Ukrainian
"ur-PK", # Urdu
"vi-VN", # Vietnamese
"cy-GB", # Welsh
]
@property
@override
def supported_formats(self) -> list[stt.AudioFormats]:
"""Return supported formats.
LiteLLM does not expose model-specific audio capabilities; the backend
may reject this.
"""
return [stt.AudioFormats.WAV]
@property
@override
def supported_codecs(self) -> list[stt.AudioCodecs]:
"""Return supported codecs.
LiteLLM does not expose model-specific audio capabilities; the backend
may reject this.
"""
return [stt.AudioCodecs.PCM]
@property
@override
def supported_bit_rates(self) -> list[stt.AudioBitRates]:
"""Return supported bit rates.
LiteLLM does not expose model-specific audio capabilities; the backend
may reject this.
"""
return [stt.AudioBitRates.BITRATE_16]
@property
@override
def supported_sample_rates(self) -> list[stt.AudioSampleRates]:
"""Return supported sample rates.
LiteLLM does not expose model-specific audio capabilities; the backend
may reject this.
"""
return [stt.AudioSampleRates.SAMPLERATE_16000]
@property
@override
def supported_channels(self) -> list[stt.AudioChannels]:
"""Return supported channels.
LiteLLM does not expose model-specific audio capabilities; the backend
may reject this.
"""
return [stt.AudioChannels.CHANNEL_MONO]
@override
async def async_process_audio_stream(
self, metadata: stt.SpeechMetadata, stream: AsyncIterable[bytes]
) -> stt.SpeechResult:
"""Process audio with the transcription endpoint."""
audio_bytes = bytearray()
async for chunk in stream:
audio_bytes.extend(chunk)
wav_buffer = io.BytesIO()
with wave.open(wav_buffer, "wb") as wav_file:
wav_file.setnchannels(metadata.channel.value)
wav_file.setsampwidth(metadata.bit_rate.value // 8)
wav_file.setframerate(metadata.sample_rate.value)
wav_file.writeframes(audio_bytes)
coordinator = self.entry.runtime_data
options = self.subentry.data
try:
prompt: str | Omit = omit
keyword_list: list[str] | Omit = omit
if options.get(CONF_STT_CUSTOM_PROMPT_KEYWORDS):
if prompt_value := options.get(CONF_STT_PROMPT):
prompt = (
Template(prompt_value, self.hass).async_render(
parse_result=False
)
or omit
)
if keywords_value := options.get(CONF_STT_KEYWORDS):
keywords = Template(keywords_value, self.hass).async_render(
parse_result=False
)
keyword_list = [
keyword.strip()
for keyword in keywords.split(",")
if keyword.strip()
] or omit
response = await coordinator.client.audio.transcriptions.create(
model=self.model,
file=("audio.wav", wav_buffer.getvalue()),
language=metadata.language.split("-")[0],
prompt=prompt,
keywords=keyword_list,
)
except TemplateError as err:
LOGGER.error("Error rendering STT template: %s", err)
except (AuthenticationError, PermissionDeniedError) as err:
await coordinator.async_request_refresh()
LOGGER.error("Authentication error during STT: %s", err)
except APIConnectionError as err:
coordinator.mark_connection_error()
LOGGER.error("Connection error during STT: %s", err)
except OpenAIError as err:
coordinator.async_set_updated_data(None)
LOGGER.error("Error during STT: %s", err)
else:
coordinator.async_set_updated_data(None)
if response.text:
return stt.SpeechResult(response.text, stt.SpeechResultState.SUCCESS)
return stt.SpeechResult(None, stt.SpeechResultState.ERROR)
+193 -2
View File
@@ -12,7 +12,13 @@ from openai import (
import pytest
from homeassistant.components.litellm.config_flow import CannotConnect, InvalidAuth
from homeassistant.components.litellm.const import CONF_PROMPT, DOMAIN
from homeassistant.components.litellm.const import (
CONF_PROMPT,
CONF_STT_CUSTOM_PROMPT_KEYWORDS,
CONF_STT_KEYWORDS,
CONF_STT_PROMPT,
DOMAIN,
)
from homeassistant.config_entries import SOURCE_USER
from homeassistant.const import CONF_API_KEY, CONF_LLM_HASS_API, CONF_MODEL, CONF_URL
from homeassistant.core import HomeAssistant
@@ -177,6 +183,189 @@ async def test_duplicate_entry(
assert result["reason"] == "already_configured"
@pytest.mark.parametrize(
("exception", "reason"),
[
(InvalidAuth(), "invalid_auth"),
(CannotConnect(), "cannot_connect"),
(Exception("unexpected"), "unknown"),
],
)
async def test_stt_subentry_exceptions(
hass: HomeAssistant,
mock_openai_client: AsyncMock,
mock_config_entry: MockConfigEntry,
exception: Exception,
reason: str,
) -> None:
"""Test STT subentry flow aborts when models cannot be fetched."""
await setup_integration(hass, mock_config_entry)
with patch(
"homeassistant.components.litellm.config_flow._get_models",
new_callable=AsyncMock,
side_effect=exception,
):
result = await hass.config_entries.subentries.async_init(
(mock_config_entry.entry_id, "stt"),
context={"source": SOURCE_USER},
)
assert result["type"] is FlowResultType.ABORT
assert result["reason"] == reason
@pytest.mark.usefixtures("mock_models")
async def test_create_stt_subentry(
hass: HomeAssistant,
mock_openai_client: AsyncMock,
mock_config_entry: MockConfigEntry,
) -> None:
"""Test creating an STT subentry."""
await setup_integration(hass, mock_config_entry)
result = await hass.config_entries.subentries.async_init(
(mock_config_entry.entry_id, "stt"),
context={"source": SOURCE_USER},
)
assert result["type"] is FlowResultType.FORM
assert result["step_id"] == "init"
result = await hass.config_entries.subentries.async_configure(
result["flow_id"],
{CONF_MODEL: "gpt-4", CONF_STT_CUSTOM_PROMPT_KEYWORDS: False},
)
assert result["type"] is FlowResultType.CREATE_ENTRY
assert result["title"] == "gpt-4"
assert result["data"] == {
CONF_MODEL: "gpt-4",
CONF_STT_CUSTOM_PROMPT_KEYWORDS: False,
}
subentry_id = get_subentry_id(mock_config_entry, "stt")
result = await mock_config_entry.start_subentry_reconfigure_flow(hass, subentry_id)
assert result["type"] is FlowResultType.FORM
result = await hass.config_entries.subentries.async_configure(
result["flow_id"],
{CONF_MODEL: "gpt-4", CONF_STT_CUSTOM_PROMPT_KEYWORDS: False},
)
assert result["type"] is FlowResultType.ABORT
assert result["reason"] == "reconfigure_successful"
async def _create_stt_with_hints(
hass: HomeAssistant, mock_config_entry: MockConfigEntry
) -> str:
"""Create an STT subentry with both optional hints enabled."""
await setup_integration(hass, mock_config_entry)
result = await hass.config_entries.subentries.async_init(
(mock_config_entry.entry_id, "stt"), context={"source": SOURCE_USER}
)
result = await hass.config_entries.subentries.async_configure(
result["flow_id"],
{CONF_MODEL: "gpt-4", CONF_STT_CUSTOM_PROMPT_KEYWORDS: True},
)
assert result["type"] is FlowResultType.FORM
result = await hass.config_entries.subentries.async_configure(
result["flow_id"],
{
CONF_MODEL: "gpt-4",
CONF_STT_CUSTOM_PROMPT_KEYWORDS: True,
CONF_STT_PROMPT: "Old prompt",
CONF_STT_KEYWORDS: "Old, keywords",
},
)
assert result["type"] is FlowResultType.CREATE_ENTRY
return get_subentry_id(mock_config_entry, "stt")
@pytest.mark.usefixtures("mock_openai_client", "mock_models")
@pytest.mark.parametrize(
("hints", "expected_hints"),
[
pytest.param({}, {}, id="clear-both-omitted"),
pytest.param(
{CONF_STT_KEYWORDS: "Old, keywords"},
{CONF_STT_KEYWORDS: "Old, keywords"},
id="clear-prompt",
),
pytest.param(
{CONF_STT_PROMPT: "Old prompt"},
{CONF_STT_PROMPT: "Old prompt"},
id="clear-keywords",
),
pytest.param(
{CONF_STT_PROMPT: "", CONF_STT_KEYWORDS: ""},
{},
id="clear-both-empty",
),
pytest.param(
{CONF_STT_PROMPT: "New prompt", CONF_STT_KEYWORDS: "New, keywords"},
{CONF_STT_PROMPT: "New prompt", CONF_STT_KEYWORDS: "New, keywords"},
id="edit-both",
),
],
)
async def test_reconfigure_stt_hints(
hass: HomeAssistant,
mock_config_entry: MockConfigEntry,
hints: dict[str, str],
expected_hints: dict[str, str],
) -> None:
"""Test saved STT hints reflect the submitted form fields."""
subentry_id = await _create_stt_with_hints(hass, mock_config_entry)
result = await mock_config_entry.start_subentry_reconfigure_flow(hass, subentry_id)
assert result["type"] is FlowResultType.FORM
result = await hass.config_entries.subentries.async_configure(
result["flow_id"],
{
CONF_MODEL: "gpt-4",
CONF_STT_CUSTOM_PROMPT_KEYWORDS: True,
**hints,
},
)
assert result["type"] is FlowResultType.ABORT
assert result["reason"] == "reconfigure_successful"
assert mock_config_entry.subentries[subentry_id].data == {
CONF_MODEL: "gpt-4",
CONF_STT_CUSTOM_PROMPT_KEYWORDS: True,
**expected_hints,
}
@pytest.mark.usefixtures("mock_openai_client", "mock_models")
async def test_reconfigure_stt_disable_hints(
hass: HomeAssistant, mock_config_entry: MockConfigEntry
) -> None:
"""Test disabling hints clears both saved STT fields."""
subentry_id = await _create_stt_with_hints(hass, mock_config_entry)
result = await mock_config_entry.start_subentry_reconfigure_flow(hass, subentry_id)
result = await hass.config_entries.subentries.async_configure(
result["flow_id"],
{CONF_MODEL: "gpt-4", CONF_STT_CUSTOM_PROMPT_KEYWORDS: False},
)
assert result["type"] is FlowResultType.FORM
result = await hass.config_entries.subentries.async_configure(
result["flow_id"],
{CONF_MODEL: "gpt-4", CONF_STT_CUSTOM_PROMPT_KEYWORDS: False},
)
assert result["type"] is FlowResultType.ABORT
assert result["reason"] == "reconfigure_successful"
assert mock_config_entry.subentries[subentry_id].data == {
CONF_MODEL: "gpt-4",
CONF_STT_CUSTOM_PROMPT_KEYWORDS: False,
}
@pytest.mark.usefixtures("mock_models")
async def test_create_conversation_agent(
hass: HomeAssistant,
@@ -369,15 +558,17 @@ async def test_reconfigure_conversation_agent_disable_llm_api(
assert key.default() == []
@pytest.mark.parametrize("subentry_type", ["conversation", "stt"])
async def test_reconfigure_entry_not_loaded(
hass: HomeAssistant,
mock_config_entry: MockConfigEntry,
subentry_type: str,
) -> None:
"""Test reconfiguring aborts when the main entry is not loaded."""
mock_config_entry.add_to_hass(hass)
result = await hass.config_entries.subentries.async_init(
(mock_config_entry.entry_id, "conversation"),
(mock_config_entry.entry_id, subentry_type),
context={"source": SOURCE_USER},
)
+291
View File
@@ -0,0 +1,291 @@
"""Test STT platform of LiteLLM integration."""
from collections.abc import AsyncIterable
import io
from unittest.mock import AsyncMock, MagicMock, patch
import wave
import httpx
from openai import (
APIConnectionError,
AuthenticationError,
OpenAIError,
PermissionDeniedError,
omit,
)
import pytest
from homeassistant.components import stt
from homeassistant.components.litellm.const import (
CONF_STT_CUSTOM_PROMPT_KEYWORDS,
CONF_STT_KEYWORDS,
CONF_STT_PROMPT,
DOMAIN,
)
from homeassistant.config_entries import ConfigSubentryData
from homeassistant.const import CONF_API_KEY, CONF_MODEL, CONF_URL
from homeassistant.core import HomeAssistant
from . import setup_integration
from .conftest import TEST_URL
from tests.common import MockConfigEntry
async def _audio_stream(*chunks: bytes) -> AsyncIterable[bytes]:
"""Yield audio chunks."""
for chunk in chunks:
yield chunk
async def _setup_stt(
hass: HomeAssistant,
subentry_options: dict[str, object] | None = None,
) -> stt.SpeechToTextEntity:
"""Set up a LiteLLM STT entity."""
entry = MockConfigEntry(
title="localhost:4000",
domain=DOMAIN,
data={CONF_URL: TEST_URL, CONF_API_KEY: "bla"},
subentries_data=[
ConfigSubentryData(
data={CONF_MODEL: "home-stt", **(subentry_options or {})},
subentry_id="STT",
subentry_type="stt",
title="home-stt",
unique_id=None,
)
],
)
await setup_integration(hass, entry)
return next(iter(hass.data[stt.DOMAIN].entities))
def _metadata() -> stt.SpeechMetadata:
"""Return the audio format supported by LiteLLM STT."""
return stt.SpeechMetadata(
language="en-US",
format=stt.AudioFormats.WAV,
codec=stt.AudioCodecs.PCM,
bit_rate=stt.AudioBitRates.BITRATE_16,
sample_rate=stt.AudioSampleRates.SAMPLERATE_16000,
channel=stt.AudioChannels.CHANNEL_MONO,
)
@pytest.mark.usefixtures("mock_openai_client")
async def test_stt_entity_properties(hass: HomeAssistant) -> None:
"""Test STT entity audio properties."""
entity = await _setup_stt(hass)
assert isinstance(entity.supported_languages, list)
assert "en-US" in entity.supported_languages
assert "pl-PL" in entity.supported_languages
assert entity.check_metadata(_metadata())
assert entity.supported_formats == [stt.AudioFormats.WAV]
assert entity.supported_codecs == [stt.AudioCodecs.PCM]
assert entity.supported_bit_rates == [stt.AudioBitRates.BITRATE_16]
assert entity.supported_sample_rates == [stt.AudioSampleRates.SAMPLERATE_16000]
assert entity.supported_channels == [stt.AudioChannels.CHANNEL_MONO]
async def test_stt(hass: HomeAssistant, mock_openai_client: AsyncMock) -> None:
"""Test transcription."""
entity = await _setup_stt(hass)
mock_openai_client.audio.transcriptions.create = AsyncMock(
return_value=MagicMock(text="Turn on the light")
)
result = await entity.async_process_audio_stream(
_metadata(), _audio_stream(b"first", b"second")
)
assert result == stt.SpeechResult(
"Turn on the light", stt.SpeechResultState.SUCCESS
)
call = mock_openai_client.audio.transcriptions.create.call_args.kwargs
assert call["model"] == "home-stt"
assert call["language"] == "en"
assert call["file"][0] == "audio.wav"
with wave.open(io.BytesIO(call["file"][1]), "rb") as wav_file:
assert wav_file.getnchannels() == 1
assert wav_file.getsampwidth() == 2
assert wav_file.getframerate() == 16000
async def test_stt_passes_prompt_and_keywords(
hass: HomeAssistant, mock_openai_client: AsyncMock
) -> None:
"""Test passing optional prompt and keyword hints to LiteLLM."""
entity = await _setup_stt(
hass,
{
CONF_STT_CUSTOM_PROMPT_KEYWORDS: True,
CONF_STT_PROMPT: "Use the configured names",
CONF_STT_KEYWORDS: "Alice, Bob, Alice",
},
)
mock_openai_client.audio.transcriptions.create = AsyncMock(
return_value=MagicMock(text="Turn on the light")
)
result = await entity.async_process_audio_stream(
_metadata(), _audio_stream(b"audio")
)
assert result.result is stt.SpeechResultState.SUCCESS
call = mock_openai_client.audio.transcriptions.create.call_args.kwargs
assert call["prompt"] == "Use the configured names"
assert call["keywords"] == ["Alice", "Bob", "Alice"]
async def test_stt_ignores_stale_hints_when_disabled(
hass: HomeAssistant, mock_openai_client: AsyncMock
) -> None:
"""Test disabled hints do not send stale configured values."""
entity = await _setup_stt(
hass,
{
CONF_STT_CUSTOM_PROMPT_KEYWORDS: False,
CONF_STT_PROMPT: "Stale prompt",
CONF_STT_KEYWORDS: "Stale, keywords",
},
)
mock_openai_client.audio.transcriptions.create = AsyncMock(
return_value=MagicMock(text="Turn on the light")
)
result = await entity.async_process_audio_stream(
_metadata(), _audio_stream(b"audio")
)
assert result.result is stt.SpeechResultState.SUCCESS
call = mock_openai_client.audio.transcriptions.create.call_args.kwargs
assert call["prompt"] is omit
assert call["keywords"] is omit
async def test_stt_renders_prompt_and_keywords_templates(
hass: HomeAssistant, mock_openai_client: AsyncMock
) -> None:
"""Test rendering optional prompt and keyword templates."""
hass.states.async_set("sensor.transcription_context", "Alice")
entity = await _setup_stt(
hass,
{
CONF_STT_CUSTOM_PROMPT_KEYWORDS: True,
CONF_STT_PROMPT: "Recognize {{ states('sensor.transcription_context') }}",
CONF_STT_KEYWORDS: "{{ states('sensor.transcription_context') }}, Bob",
},
)
mock_openai_client.audio.transcriptions.create = AsyncMock(
return_value=MagicMock(text="Turn on the light")
)
result = await entity.async_process_audio_stream(
_metadata(), _audio_stream(b"audio")
)
assert result.result is stt.SpeechResultState.SUCCESS
call = mock_openai_client.audio.transcriptions.create.call_args.kwargs
assert call["prompt"] == "Recognize Alice"
assert call["keywords"] == ["Alice", "Bob"]
async def test_stt_template_error(
hass: HomeAssistant, mock_openai_client: AsyncMock
) -> None:
"""Test a template error returns an STT error."""
entity = await _setup_stt(
hass, {CONF_STT_CUSTOM_PROMPT_KEYWORDS: True, CONF_STT_PROMPT: "{{ 1 / 0 }}"}
)
result = await entity.async_process_audio_stream(
_metadata(), _audio_stream(b"audio")
)
assert result == stt.SpeechResult(None, stt.SpeechResultState.ERROR)
mock_openai_client.audio.transcriptions.create.assert_not_called()
async def test_stt_empty_response_keeps_entity_available(
hass: HomeAssistant, mock_openai_client: AsyncMock
) -> None:
"""Test an empty response keeps the proxy available."""
entity = await _setup_stt(hass)
coordinator = entity.entry.runtime_data
coordinator.mark_connection_error()
mock_openai_client.audio.transcriptions.create = AsyncMock(
return_value=MagicMock(text="")
)
result = await entity.async_process_audio_stream(
_metadata(), _audio_stream(b"audio")
)
assert result == stt.SpeechResult(None, stt.SpeechResultState.ERROR)
assert entity.available
async def test_stt_connection_error_marks_entity_unavailable(
hass: HomeAssistant, mock_openai_client: AsyncMock
) -> None:
"""Test connection errors mark the entity unavailable."""
entity = await _setup_stt(hass)
mock_openai_client.audio.transcriptions.create = AsyncMock(
side_effect=APIConnectionError(request=None)
)
result = await entity.async_process_audio_stream(
_metadata(), _audio_stream(b"audio")
)
assert result.result is stt.SpeechResultState.ERROR
assert not entity.available
@pytest.mark.parametrize("error_cls", [AuthenticationError, PermissionDeniedError])
async def test_stt_auth_error_refreshes_coordinator(
hass: HomeAssistant,
mock_openai_client: AsyncMock,
error_cls: type[AuthenticationError | PermissionDeniedError],
) -> None:
"""Test authentication errors refresh coordinator state."""
entity = await _setup_stt(hass)
mock_openai_client.audio.transcriptions.create = AsyncMock(
side_effect=error_cls(
message="invalid api key",
response=httpx.Response(
401, request=httpx.Request("POST", "http://localhost")
),
body=None,
)
)
coordinator = entity.entry.runtime_data
with patch.object(
coordinator, "async_request_refresh", new_callable=AsyncMock
) as refresh:
result = await entity.async_process_audio_stream(
_metadata(), _audio_stream(b"audio")
)
assert result.result is stt.SpeechResultState.ERROR
refresh.assert_awaited_once()
async def test_stt_provider_error_keeps_entity_available(
hass: HomeAssistant, mock_openai_client: AsyncMock
) -> None:
"""Test provider errors do not mark the proxy unavailable."""
entity = await _setup_stt(hass)
mock_openai_client.audio.transcriptions.create = AsyncMock(
side_effect=OpenAIError("bad request")
)
result = await entity.async_process_audio_stream(
_metadata(), _audio_stream(b"audio")
)
assert result.result is stt.SpeechResultState.ERROR
assert entity.available