From a88f859aaecb9716a9302b19ebd6bea12b2618a4 Mon Sep 17 00:00:00 2001 From: Marcin Juraszek <10106352+korasino@users.noreply.github.com> Date: Sat, 3 Oct 2026 22:17:48 +0200 Subject: [PATCH] Add LiteLLM speech-to-text support (#183073) --- homeassistant/components/litellm/__init__.py | 2 +- .../components/litellm/config_flow.py | 121 +++++++- homeassistant/components/litellm/const.py | 3 + homeassistant/components/litellm/strings.json | 30 ++ homeassistant/components/litellm/stt.py | 232 ++++++++++++++ tests/components/litellm/test_config_flow.py | 195 +++++++++++- tests/components/litellm/test_stt.py | 291 ++++++++++++++++++ 7 files changed, 870 insertions(+), 4 deletions(-) create mode 100644 homeassistant/components/litellm/stt.py create mode 100644 tests/components/litellm/test_stt.py diff --git a/homeassistant/components/litellm/__init__.py b/homeassistant/components/litellm/__init__.py index 5447eeb01454..17687e882258 100644 --- a/homeassistant/components/litellm/__init__.py +++ b/homeassistant/components/litellm/__init__.py @@ -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: diff --git a/homeassistant/components/litellm/config_flow.py b/homeassistant/components/litellm/config_flow.py index e946cdb310b7..148c12bff88a 100644 --- a/homeassistant/components/litellm/config_flow.py +++ b/homeassistant/components/litellm/config_flow.py @@ -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), + ) diff --git a/homeassistant/components/litellm/const.py b/homeassistant/components/litellm/const.py index 8f645e234519..f546e47e08f8 100644 --- a/homeassistant/components/litellm/const.py +++ b/homeassistant/components/litellm/const.py @@ -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. diff --git a/homeassistant/components/litellm/strings.json b/homeassistant/components/litellm/strings.json index 6e03c07bbebc..ee67d48ce800 100644 --- a/homeassistant/components/litellm/strings.json +++ b/homeassistant/components/litellm/strings.json @@ -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." + } + } } } } diff --git a/homeassistant/components/litellm/stt.py b/homeassistant/components/litellm/stt.py new file mode 100644 index 000000000000..f6e3fefbbd12 --- /dev/null +++ b/homeassistant/components/litellm/stt.py @@ -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) diff --git a/tests/components/litellm/test_config_flow.py b/tests/components/litellm/test_config_flow.py index 57edc21d4382..dcf2296efb4b 100644 --- a/tests/components/litellm/test_config_flow.py +++ b/tests/components/litellm/test_config_flow.py @@ -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}, ) diff --git a/tests/components/litellm/test_stt.py b/tests/components/litellm/test_stt.py new file mode 100644 index 000000000000..fd13fc9373ea --- /dev/null +++ b/tests/components/litellm/test_stt.py @@ -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