import logging import base64 from functools import partial from collections.abc import AsyncIterable from homeassistant.components import stt from homeassistant.config_entries import ConfigEntry from homeassistant.core import HomeAssistant from homeassistant.helpers.entity_platform import AddEntitiesCallback from .tencentcloud_api import TencentCloudAsrAPi from .tencentcloud_api import ModuleSupportLanguage from .tencentcloud_api import DefaultModel from . import ( ConvertNumModeKey, FilterDirtyKey, FilterModalKey, FilterPuncKey, ModelKey, SecretIdKey, SecretKeyKey, ) _LOGGER = logging.getLogger(__name__) async def async_setup_entry( hass: HomeAssistant, config_entry: ConfigEntry, async_add_entities: AddEntitiesCallback, ) -> None: async_add_entities([ASRSTT(hass, config_entry)]) class ASRSTT(stt.SpeechToTextEntity): def __init__(self, hass: HomeAssistant, config_entry: ConfigEntry) -> None: secretId = config_entry.data.get(SecretIdKey, "") secretKey = config_entry.data.get(SecretKeyKey, "") self.tencentCloudApi = TencentCloudAsrAPi(secretId, secretKey) self._attr_unique_id = f"tencentcloud_asr_stt" self._attr_name = f"TencentCloud Asr STT" self.hass = hass self.config_entry = config_entry @property def supported_languages(self) -> list[str]: model: str = self.hass.data.get(ModelKey, DefaultModel) return ModuleSupportLanguage.get(model, ["zh"]) @property def supported_formats(self) -> list[stt.AudioFormats]: return [stt.AudioFormats.WAV] @property def supported_codecs(self) -> list[stt.AudioCodecs]: return [stt.AudioCodecs.PCM] @property def supported_bit_rates(self) -> list[stt.AudioBitRates]: return [stt.AudioBitRates.BITRATE_16] @property def supported_sample_rates(self) -> list[stt.AudioSampleRates]: model: str = self.hass.data.get(ModelKey, DefaultModel) if model.startswith("8k_"): return [stt.AudioSampleRates.SAMPLERATE_8000] return [stt.AudioSampleRates.SAMPLERATE_16000] @property def supported_channels(self) -> list[stt.AudioChannels]: return [stt.AudioChannels.CHANNEL_MONO] async def async_process_audio_stream( self, metadata: stt.SpeechMetadata, stream: AsyncIterable[bytes] ) -> stt.SpeechResult: _LOGGER.debug("process_audio_stream start") audio = b"" async for chunk in stream: audio += chunk data = base64.b64encode(audio).decode("utf8", "ignore").strip() config = {**self.config_entry.data, **self.config_entry.options} model: str = config.get(ModelKey, DefaultModel) _LOGGER.debug( "process_audio_stream transcribe: audio_bytes=%s model=%s", len(audio), model, ) ret, result = await self.hass.async_add_executor_job( partial( self.tencentCloudApi.SentenceRecognition, model, data, len(audio), config.get(FilterDirtyKey, "0"), config.get(FilterModalKey, False), config.get(FilterPuncKey, False), config.get(ConvertNumModeKey, True), ) ) if not ret: return stt.SpeechResult(result, stt.SpeechResultState.ERROR) if not result: return stt.SpeechResult("未识别到有效语音", stt.SpeechResultState.SUCCESS) return stt.SpeechResult(result, stt.SpeechResultState.SUCCESS)