mirror of
https://github.com/home-assistant/core.git
synced 2026-10-07 06:50:41 -04:00
Add LiteLLM speech-to-text support (#183073)
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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."
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
@@ -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},
|
||||
)
|
||||
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user