From cc7fe8626e61a6437d1cbf53a1655fadc07a4353 Mon Sep 17 00:00:00 2001 From: xyzm Date: Fri, 6 Feb 2026 09:21:49 +0800 Subject: [PATCH] =?UTF-8?q?```=20feat(stt):=20=E6=96=B0=E5=A2=9E=E8=B1=86?= =?UTF-8?q?=E5=8C=85=E8=AF=AD=E9=9F=B3=E8=AF=86=E5=88=AB=E9=9B=86=E6=88=90?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 实现完整的 Doubao STT 组件,支持中文语音识别 - 提供 UI 配置界面,无需手动编辑配置文件 - 支持实时流式识别和自动标点符号添加 - 集成自动设备注册和凭据管理功能 - 添加 HACS 和 hassfest 验证工作流 - 配置 .gitignore 和 Apache 2.0 许可证文件 - 完善 README 文档和故障排除指南 ``` --- .github/workflows/validate.yml | 29 + .gitignore | 3 + LICENSE | 201 ++++++ README.md | 159 +++++ custom_components/hass_stt_doubao/__init__.py | 55 ++ .../hass_stt_doubao/config_flow.py | 171 +++++ custom_components/hass_stt_doubao/const.py | 26 + .../hass_stt_doubao/doubaoime_asr/__init__.py | 23 + .../hass_stt_doubao/doubaoime_asr/asr.proto | 75 ++ .../hass_stt_doubao/doubaoime_asr/asr.py | 656 ++++++++++++++++++ .../hass_stt_doubao/doubaoime_asr/asr_pb2.py | 40 ++ .../hass_stt_doubao/doubaoime_asr/asr_pb2.pyi | 55 ++ .../hass_stt_doubao/doubaoime_asr/audio.py | 77 ++ .../hass_stt_doubao/doubaoime_asr/config.py | 231 ++++++ .../doubaoime_asr/constants.py | 43 ++ .../hass_stt_doubao/doubaoime_asr/device.py | 321 +++++++++ .../hass_stt_doubao/manifest.json | 18 + .../hass_stt_doubao/requirements.txt | 6 + .../hass_stt_doubao/strings.json | 33 + custom_components/hass_stt_doubao/stt.py | 162 +++++ .../hass_stt_doubao/translations/zh.json | 33 + hacs.json | 6 + 22 files changed, 2423 insertions(+) create mode 100644 .github/workflows/validate.yml create mode 100644 .gitignore create mode 100644 LICENSE create mode 100644 README.md create mode 100644 custom_components/hass_stt_doubao/__init__.py create mode 100644 custom_components/hass_stt_doubao/config_flow.py create mode 100644 custom_components/hass_stt_doubao/const.py create mode 100644 custom_components/hass_stt_doubao/doubaoime_asr/__init__.py create mode 100644 custom_components/hass_stt_doubao/doubaoime_asr/asr.proto create mode 100644 custom_components/hass_stt_doubao/doubaoime_asr/asr.py create mode 100644 custom_components/hass_stt_doubao/doubaoime_asr/asr_pb2.py create mode 100644 custom_components/hass_stt_doubao/doubaoime_asr/asr_pb2.pyi create mode 100644 custom_components/hass_stt_doubao/doubaoime_asr/audio.py create mode 100644 custom_components/hass_stt_doubao/doubaoime_asr/config.py create mode 100644 custom_components/hass_stt_doubao/doubaoime_asr/constants.py create mode 100644 custom_components/hass_stt_doubao/doubaoime_asr/device.py create mode 100644 custom_components/hass_stt_doubao/manifest.json create mode 100644 custom_components/hass_stt_doubao/requirements.txt create mode 100644 custom_components/hass_stt_doubao/strings.json create mode 100644 custom_components/hass_stt_doubao/stt.py create mode 100644 custom_components/hass_stt_doubao/translations/zh.json create mode 100644 hacs.json diff --git a/.github/workflows/validate.yml b/.github/workflows/validate.yml new file mode 100644 index 0000000..2b47432 --- /dev/null +++ b/.github/workflows/validate.yml @@ -0,0 +1,29 @@ +name: "Validation And Formatting" +on: + push: + pull_request: + workflow_dispatch: + schedule: + - cron: '0 0 * * *' +jobs: + hacs: + runs-on: "ubuntu-latest" + name: HACS + steps: + - name: Check out the repository + uses: "actions/checkout@v3" + + - name: HACS validation + uses: "hacs/action@main" + with: + category: "integration" + + hassfest: + runs-on: "ubuntu-latest" + name: Hassfest + steps: + - name: Check out the repository + uses: "actions/checkout@v3" + + - name: Hassfest validation + uses: "home-assistant/actions/hassfest@master" diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..4b599c5 --- /dev/null +++ b/.gitignore @@ -0,0 +1,3 @@ +.idea/ +*.mp3 +*.pyc \ No newline at end of file diff --git a/LICENSE b/LICENSE new file mode 100644 index 0000000..261eeb9 --- /dev/null +++ b/LICENSE @@ -0,0 +1,201 @@ + Apache License + Version 2.0, January 2004 + http://www.apache.org/licenses/ + + TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION + + 1. Definitions. + + "License" shall mean the terms and conditions for use, reproduction, + and distribution as defined by Sections 1 through 9 of this document. + + "Licensor" shall mean the copyright owner or entity authorized by + the copyright owner that is granting the License. + + "Legal Entity" shall mean the union of the acting entity and all + other entities that control, are controlled by, or are under common + control with that entity. For the purposes of this definition, + "control" means (i) the power, direct or indirect, to cause the + direction or management of such entity, whether by contract or + otherwise, or (ii) ownership of fifty percent (50%) or more of the + outstanding shares, or (iii) beneficial ownership of such entity. + + "You" (or "Your") shall mean an individual or Legal Entity + exercising permissions granted by this License. + + "Source" form shall mean the preferred form for making modifications, + including but not limited to software source code, documentation + source, and configuration files. + + "Object" form shall mean any form resulting from mechanical + transformation or translation of a Source form, including but + not limited to compiled object code, generated documentation, + and conversions to other media types. + + "Work" shall mean the work of authorship, whether in Source or + Object form, made available under the License, as indicated by a + copyright notice that is included in or attached to the work + (an example is provided in the Appendix below). + + "Derivative Works" shall mean any work, whether in Source or Object + form, that is based on (or derived from) the Work and for which the + editorial revisions, annotations, elaborations, or other modifications + represent, as a whole, an original work of authorship. For the purposes + of this License, Derivative Works shall not include works that remain + separable from, or merely link (or bind by name) to the interfaces of, + the Work and Derivative Works thereof. + + "Contribution" shall mean any work of authorship, including + the original version of the Work and any modifications or additions + to that Work or Derivative Works thereof, that is intentionally + submitted to Licensor for inclusion in the Work by the copyright owner + or by an individual or Legal Entity authorized to submit on behalf of + the copyright owner. For the purposes of this definition, "submitted" + means any form of electronic, verbal, or written communication sent + to the Licensor or its representatives, including but not limited to + communication on electronic mailing lists, source code control systems, + and issue tracking systems that are managed by, or on behalf of, the + Licensor for the purpose of discussing and improving the Work, but + excluding communication that is conspicuously marked or otherwise + designated in writing by the copyright owner as "Not a Contribution." + + "Contributor" shall mean Licensor and any individual or Legal Entity + on behalf of whom a Contribution has been received by Licensor and + subsequently incorporated within the Work. + + 2. Grant of Copyright License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + copyright license to reproduce, prepare Derivative Works of, + publicly display, publicly perform, sublicense, and distribute the + Work and such Derivative Works in Source or Object form. + + 3. Grant of Patent License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + (except as stated in this section) patent license to make, have made, + use, offer to sell, sell, import, and otherwise transfer the Work, + where such license applies only to those patent claims licensable + by such Contributor that are necessarily infringed by their + Contribution(s) alone or by combination of their Contribution(s) + with the Work to which such Contribution(s) was submitted. If You + institute patent litigation against any entity (including a + cross-claim or counterclaim in a lawsuit) alleging that the Work + or a Contribution incorporated within the Work constitutes direct + or contributory patent infringement, then any patent licenses + granted to You under this License for that Work shall terminate + as of the date such litigation is filed. + + 4. Redistribution. You may reproduce and distribute copies of the + Work or Derivative Works thereof in any medium, with or without + modifications, and in Source or Object form, provided that You + meet the following conditions: + + (a) You must give any other recipients of the Work or + Derivative Works a copy of this License; and + + (b) You must cause any modified files to carry prominent notices + stating that You changed the files; and + + (c) You must retain, in the Source form of any Derivative Works + that You distribute, all copyright, patent, trademark, and + attribution notices from the Source form of the Work, + excluding those notices that do not pertain to any part of + the Derivative Works; and + + (d) If the Work includes a "NOTICE" text file as part of its + distribution, then any Derivative Works that You distribute must + include a readable copy of the attribution notices contained + within such NOTICE file, excluding those notices that do not + pertain to any part of the Derivative Works, in at least one + of the following places: within a NOTICE text file distributed + as part of the Derivative Works; within the Source form or + documentation, if provided along with the Derivative Works; or, + within a display generated by the Derivative Works, if and + wherever such third-party notices normally appear. The contents + of the NOTICE file are for informational purposes only and + do not modify the License. You may add Your own attribution + notices within Derivative Works that You distribute, alongside + or as an addendum to the NOTICE text from the Work, provided + that such additional attribution notices cannot be construed + as modifying the License. + + You may add Your own copyright statement to Your modifications and + may provide additional or different license terms and conditions + for use, reproduction, or distribution of Your modifications, or + for any such Derivative Works as a whole, provided Your use, + reproduction, and distribution of the Work otherwise complies with + the conditions stated in this License. + + 5. Submission of Contributions. Unless You explicitly state otherwise, + any Contribution intentionally submitted for inclusion in the Work + by You to the Licensor shall be under the terms and conditions of + this License, without any additional terms or conditions. + Notwithstanding the above, nothing herein shall supersede or modify + the terms of any separate license agreement you may have executed + with Licensor regarding such Contributions. + + 6. Trademarks. This License does not grant permission to use the trade + names, trademarks, service marks, or product names of the Licensor, + except as required for reasonable and customary use in describing the + origin of the Work and reproducing the content of the NOTICE file. + + 7. Disclaimer of Warranty. Unless required by applicable law or + agreed to in writing, Licensor provides the Work (and each + Contributor provides its Contributions) on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or + implied, including, without limitation, any warranties or conditions + of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A + PARTICULAR PURPOSE. You are solely responsible for determining the + appropriateness of using or redistributing the Work and assume any + risks associated with Your exercise of permissions under this License. + + 8. Limitation of Liability. In no event and under no legal theory, + whether in tort (including negligence), contract, or otherwise, + unless required by applicable law (such as deliberate and grossly + negligent acts) or agreed to in writing, shall any Contributor be + liable to You for damages, including any direct, indirect, special, + incidental, or consequential damages of any character arising as a + result of this License or out of the use or inability to use the + Work (including but not limited to damages for loss of goodwill, + work stoppage, computer failure or malfunction, or any and all + other commercial damages or losses), even if such Contributor + has been advised of the possibility of such damages. + + 9. Accepting Warranty or Additional Liability. While redistributing + the Work or Derivative Works thereof, You may choose to offer, + and charge a fee for, acceptance of support, warranty, indemnity, + or other liability obligations and/or rights consistent with this + License. However, in accepting such obligations, You may act only + on Your own behalf and on Your sole responsibility, not on behalf + of any other Contributor, and only if You agree to indemnify, + defend, and hold each Contributor harmless for any liability + incurred by, or claims asserted against, such Contributor by reason + of your accepting any such warranty or additional liability. + + END OF TERMS AND CONDITIONS + + APPENDIX: How to apply the Apache License to your work. + + To apply the Apache License to your work, attach the following + boilerplate notice, with the fields enclosed by brackets "[]" + replaced with your own identifying information. (Don't include + the brackets!) The text should be enclosed in the appropriate + comment syntax for the file format. We also recommend that a + file or class name and description of purpose be included on the + same "printed page" as the copyright notice for easier + identification within third-party archives. + + Copyright [yyyy] [name of copyright owner] + + Licensed under the Apache License, Version 2.0 (the "License"); + you may not use this file except in compliance with the License. + You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. diff --git a/README.md b/README.md new file mode 100644 index 0000000..7b98ccc --- /dev/null +++ b/README.md @@ -0,0 +1,159 @@ +# Doubao Speech-to-Text for Home Assistant + +豆包语音识别 Home Assistant 集成插件,基于 [doubaoime-asr](https://github.com/starccy/doubaoime-asr) 开发。 + + +## 功能特性 + +- ✅ 完全兼容 Home Assistant STT 组件规范 +- ✅ UI 配置界面,无需手动编辑配置文件 +- ✅ 自动设备注册和凭据管理 +- ✅ 支持中文语音识别(zh-CN, zh) +- ✅ 实时流式识别 +- ✅ 自动标点符号添加 +- ✅ 保留原有设备伪装功能 + +## 系统要求 + +- Home Assistant 2024.1+ +- Python 3.11+ + + +## 安装方法 + +### 方法 1: 手动安装 + +1. 将整个 `custom_components/hass_stt_doubao` 目录复制到 Home Assistant 配置目录的 `custom_components` 文件夹下: + +2. 重启 Home Assistant + +### 方法 2: HACS 安装(推荐) + +1. 确保你已经在 Home Assistant 中安装并配置了 HACS(Home Assistant Community Store)。 +2. 点击 HACS 界面左上角的菜单,选择 自定义仓库(Custom repositories),在弹出的窗口中输入仓库地址 https://github.com/xyzmos/hass_stt_doubao,类别选择 集成,然后点击 添加。之后再搜索 "Doubao STT" 并安装。 +3. 重启 Home Assistant + +## 配置 + +### 通过 UI 配置(推荐) + +1. 进入 Home Assistant 设置 → 设备与服务 +2. 点击右下角 "添加集成" +3. 搜索 "Doubao Speech to Text" +4. 按照向导完成配置: + - **凭据文件路径**:默认 `doubao_credentials.json`(相对于 HA 配置目录) + - **启用标点符号**:默认启用 + +### 配置说明 + +- **凭据文件路径**:首次运行时会自动注册虚拟设备并保存凭据到此文件,避免重复注册 +- **启用标点符号**:是否在识别结果中自动添加标点符号 + +## 使用示例 + +### 在 Assist 中使用 + +配置完成后,Doubao STT 会自动出现在 Assist 管道配置中: + +1. 进入设置 → Assist → 管道 +2. 选择或创建一个管道 +3. 在 "语音转文字" 选项中选择 "Doubao STT" +4. 保存配置 + +### 在自动化中使用 + +```yaml +automation: + - alias: "语音控制示例" + trigger: + - platform: event + event_type: voice_command + action: + - service: stt.process + data: + entity_id: stt.doubao_stt + language: zh-CN +``` + +## 技术架构 + +``` +Home Assistant Audio Stream (PCM 16kHz Mono) + ↓ + DoubaoSTTEntity + ↓ + Audio Encoder (PCM → Opus) + ↓ + DoubaoASR WebSocket Client + ↓ + Doubao ASR Service + ↓ + Recognition Result (Text) +``` + +## 目录结构 + +``` +custom_components/hass_stt_doubao/ +├── __init__.py # 集成初始化 +├── manifest.json # 集成元数据 +├── const.py # 常量定义 +├── config_flow.py # 配置流 +├── stt.py # STT 实体实现 +├── strings.json # UI 字符串 +├── translations/ +│ └── zh.json # 中文翻译 +└── doubaoime_asr/ # doubaoime-asr 核心模块 + ├── __init__.py + ├── asr.py + ├── audio.py + ├── config.py + ├── constants.py + ├── device.py + └── asr_pb2.py +``` + +## 故障排除 + +### 无法连接到 Doubao 服务 + +1. 检查网络连接 +2. 确认系统已安装 libopus0 +3. 查看 Home Assistant 日志:`设置 → 系统 → 日志` + +### 识别结果为空 + +1. 确保音频格式正确(PCM 16kHz Mono) +2. 检查麦克风是否正常工作 +3. 尝试重新配置集成 + +### 凭据失效 + +1. 删除凭据文件(默认在配置目录下的 `doubao_credentials.json`) +2. 重新配置集成,会自动重新注册设备 + +## 免责声明 + +本项目基于 [doubaoime-asr](https://github.com/starccy/doubaoime-asr) - 核心 ASR 的实现,**非官方提供的 API**。 + +- 本项目仅供学习和研究目的 +- 不保证未来的可用性和稳定性 +- 服务端协议可能随时变更导致功能失效 + +## 许可证 + +MIT License + +## 致谢 + +- [doubaoime-asr](https://github.com/starccy/doubaoime-asr) - 核心 ASR 实现 +- [Home Assistant](https://www.home-assistant.io/) - 智能家居平台 + +## 更新日志 + +### v1.0.0 (2026-02-05) +- 初始版本 +- 完整的 STT 组件实现 +- UI 配置支持 +- 自动设备注册 +- 中文语音识别 diff --git a/custom_components/hass_stt_doubao/__init__.py b/custom_components/hass_stt_doubao/__init__.py new file mode 100644 index 0000000..690cf57 --- /dev/null +++ b/custom_components/hass_stt_doubao/__init__.py @@ -0,0 +1,55 @@ +"""The Doubao Speech-to-Text integration.""" +from __future__ import annotations + +import logging +from pathlib import Path + +from homeassistant.config_entries import ConfigEntry +from homeassistant.const import Platform +from homeassistant.core import HomeAssistant + +from .const import ( + DOMAIN, + CONF_CREDENTIAL_PATH, + CONF_ENABLE_PUNCTUATION, + DEFAULT_CREDENTIAL_PATH, + DEFAULT_ENABLE_PUNCTUATION, +) + +_LOGGER = logging.getLogger(__name__) + +PLATFORMS: list[Platform] = [Platform.STT] + + +async def async_setup_entry(hass: HomeAssistant, entry: ConfigEntry) -> bool: + """Set up Doubao Speech-to-Text from a config entry.""" + hass.data.setdefault(DOMAIN, {}) + + # 获取配置 + credential_path = entry.data.get(CONF_CREDENTIAL_PATH, DEFAULT_CREDENTIAL_PATH) + enable_punctuation = entry.data.get(CONF_ENABLE_PUNCTUATION, DEFAULT_ENABLE_PUNCTUATION) + + # 将相对路径转换为绝对路径 + if not Path(credential_path).is_absolute(): + credential_path = hass.config.path(credential_path) + + # 存储配置供 STT 实体使用 + hass.data[DOMAIN][entry.entry_id] = { + CONF_CREDENTIAL_PATH: credential_path, + CONF_ENABLE_PUNCTUATION: enable_punctuation, + } + + # 加载 STT 平台 + await hass.config_entries.async_forward_entry_setups(entry, PLATFORMS) + + _LOGGER.info("Doubao STT 集成已设置完成") + + return True + + +async def async_unload_entry(hass: HomeAssistant, entry: ConfigEntry) -> bool: + """Unload a config entry.""" + if unload_ok := await hass.config_entries.async_unload_platforms(entry, PLATFORMS): + hass.data[DOMAIN].pop(entry.entry_id) + + return unload_ok diff --git a/custom_components/hass_stt_doubao/config_flow.py b/custom_components/hass_stt_doubao/config_flow.py new file mode 100644 index 0000000..48599ef --- /dev/null +++ b/custom_components/hass_stt_doubao/config_flow.py @@ -0,0 +1,171 @@ +"""Config flow for Doubao Speech-to-Text integration.""" +from __future__ import annotations + +import logging +from pathlib import Path +from typing import Any + +import voluptuous as vol + +from homeassistant import config_entries +from homeassistant.core import HomeAssistant, callback +from homeassistant.data_entry_flow import FlowResult +from homeassistant.exceptions import HomeAssistantError + +from .const import ( + DOMAIN, + CONF_CREDENTIAL_PATH, + CONF_ENABLE_PUNCTUATION, + DEFAULT_CREDENTIAL_PATH, + DEFAULT_ENABLE_PUNCTUATION, +) + +_LOGGER = logging.getLogger(__name__) + + +async def validate_input(hass: HomeAssistant, data: dict[str, Any]) -> dict[str, Any]: + """Validate the user input allows us to connect. + + Data has the keys from STEP_USER_DATA_SCHEMA with values provided by the user. + """ + from .doubaoime_asr import ASRConfig, DoubaoASR, ASRError + + credential_path = data.get(CONF_CREDENTIAL_PATH, DEFAULT_CREDENTIAL_PATH) + enable_punctuation = data.get(CONF_ENABLE_PUNCTUATION, DEFAULT_ENABLE_PUNCTUATION) + + # 将相对路径转换为 Home Assistant 配置目录下的绝对路径 + if not Path(credential_path).is_absolute(): + credential_path = hass.config.path(credential_path) + + # 尝试初始化配置并验证凭据 + try: + config = ASRConfig( + credential_path=credential_path, + enable_punctuation=enable_punctuation, + ) + # 确保凭据已初始化(会自动注册设备如果需要) + await config.async_ensure_credentials() + + # 简单验证:检查是否成功获取了 device_id 和 token + if not config.device_id or not config.token: + raise CannotConnect("无法获取设备凭据") + + except ASRError as err: + _LOGGER.error("验证 Doubao STT 配置失败: %s", err) + raise CannotConnect(str(err)) from err + except Exception as err: + _LOGGER.exception("验证 Doubao STT 配置时发生未知错误") + raise CannotConnect(str(err)) from err + + # 返回用户可读的标题信息 + return { + "title": "Doubao STT", + "credential_path": credential_path, + } + + +class ConfigFlow(config_entries.ConfigFlow, domain=DOMAIN): + """Handle a config flow for Doubao Speech-to-Text.""" + + VERSION = 1 + + async def async_step_user( + self, user_input: dict[str, Any] | None = None + ) -> FlowResult: + """Handle the initial step.""" + errors: dict[str, str] = {} + + if user_input is not None: + try: + info = await validate_input(self.hass, user_input) + except CannotConnect: + errors["base"] = "cannot_connect" + except Exception: # pylint: disable=broad-except + _LOGGER.exception("Unexpected exception") + errors["base"] = "unknown" + else: + return self.async_create_entry( + title=info["title"], + data=user_input, + ) + + # 显示配置表单 + data_schema = vol.Schema( + { + vol.Optional( + CONF_CREDENTIAL_PATH, + default=DEFAULT_CREDENTIAL_PATH, + ): str, + vol.Optional( + CONF_ENABLE_PUNCTUATION, + default=DEFAULT_ENABLE_PUNCTUATION, + ): bool, + } + ) + + return self.async_show_form( + step_id="user", + data_schema=data_schema, + errors=errors, + ) + + @staticmethod + @callback + def async_get_options_flow( + config_entry: config_entries.ConfigEntry, + ) -> OptionsFlowHandler: + """Get the options flow for this handler.""" + return OptionsFlowHandler() + + +class OptionsFlowHandler(config_entries.OptionsFlow): + """Handle options flow for Doubao STT.""" + + async def async_step_init( + self, user_input: dict[str, Any] | None = None + ) -> FlowResult: + """Manage the options.""" + errors: dict[str, str] = {} + + if user_input is not None: + # 验证新的配置 + try: + await validate_input(self.hass, user_input) + except CannotConnect: + errors["base"] = "cannot_connect" + except Exception: # pylint: disable=broad-except + _LOGGER.exception("Unexpected exception") + errors["base"] = "unknown" + else: + return self.async_create_entry(title="", data=user_input) + + # 获取当前配置值 + current_credential_path = self.config_entry.data.get( + CONF_CREDENTIAL_PATH, DEFAULT_CREDENTIAL_PATH + ) + current_enable_punctuation = self.config_entry.data.get( + CONF_ENABLE_PUNCTUATION, DEFAULT_ENABLE_PUNCTUATION + ) + + data_schema = vol.Schema( + { + vol.Optional( + CONF_CREDENTIAL_PATH, + default=current_credential_path, + ): str, + vol.Optional( + CONF_ENABLE_PUNCTUATION, + default=current_enable_punctuation, + ): bool, + } + ) + + return self.async_show_form( + step_id="init", + data_schema=data_schema, + errors=errors, + ) + + +class CannotConnect(HomeAssistantError): + """Error to indicate we cannot connect.""" diff --git a/custom_components/hass_stt_doubao/const.py b/custom_components/hass_stt_doubao/const.py new file mode 100644 index 0000000..93c9e71 --- /dev/null +++ b/custom_components/hass_stt_doubao/const.py @@ -0,0 +1,26 @@ +"""Constants for the Doubao Speech-to-Text integration.""" + +from typing import Final + +# Integration domain +DOMAIN: Final = "hass_stt_doubao" + +# Configuration keys +CONF_CREDENTIAL_PATH: Final = "credential_path" +CONF_ENABLE_PUNCTUATION: Final = "enable_punctuation" + +# Default values +DEFAULT_CREDENTIAL_PATH: Final = "doubao_credentials.json" +DEFAULT_ENABLE_PUNCTUATION: Final = True +DEFAULT_SAMPLE_RATE: Final = 16000 +DEFAULT_CHANNELS: Final = 1 +DEFAULT_APP_NAME: Final = "com.android.chrome" + +# Supported languages +SUPPORTED_LANGUAGES: Final = ["zh-CN", "zh"] + +# Service names (if needed) +SERVICE_TRANSCRIBE: Final = "transcribe" + +# Event types +EVENT_DOUBAO_STT_RESULT: Final = f"{DOMAIN}_result" diff --git a/custom_components/hass_stt_doubao/doubaoime_asr/__init__.py b/custom_components/hass_stt_doubao/doubaoime_asr/__init__.py new file mode 100644 index 0000000..4adbd21 --- /dev/null +++ b/custom_components/hass_stt_doubao/doubaoime_asr/__init__.py @@ -0,0 +1,23 @@ +from .asr import ( + DoubaoASR, + ASRResponse, + ASRError, + ResponseType, + AudioChunk, + transcribe, + transcribe_stream, + transcribe_realtime, +) +from .config import ASRConfig + +__all__ = [ + "DoubaoASR", + "ASRConfig", + "ASRResponse", + "ASRError", + "ResponseType", + "AudioChunk", + "transcribe", + "transcribe_stream", + "transcribe_realtime", +] diff --git a/custom_components/hass_stt_doubao/doubaoime_asr/asr.proto b/custom_components/hass_stt_doubao/doubaoime_asr/asr.proto new file mode 100644 index 0000000..d973486 --- /dev/null +++ b/custom_components/hass_stt_doubao/doubaoime_asr/asr.proto @@ -0,0 +1,75 @@ +syntax = "proto3"; + +package asr; + +// ASR WebSocket 请求消息 +// 用于 StartTask, StartSession, FinishSession, TaskRequest 等操作 +message AsrRequest { + // 认证 token (从 Settings API 获取的 app_key) + string token = 2; + + // 服务名称,固定为 "ASR" + string service_name = 3; + + // 方法名称 + // - "StartTask": 启动任务 + // - "StartSession": 启动会话 + // - "FinishSession": 结束会话 + // - "TaskRequest": 发送音频数据 + string method_name = 5; + + // JSON 格式的负载数据 + // - StartSession: 会话配置 (audio_info, enable_punctuation 等) + // - TaskRequest: 元数据 {"extra": {}, "timestamp_ms": xxx} + string payload = 6; + + // 音频数据 (仅 TaskRequest 使用) + bytes audio_data = 7; + + // 请求 ID (UUID) + string request_id = 8; + + // 帧状态 (仅 TaskRequest 使用) + FrameState frame_state = 9; +} + +// ASR WebSocket 响应消息 +message AsrResponse { + // 请求 ID + string request_id = 1; + + // 任务 ID + string task_id = 2; + + // 服务名称 + string service_name = 3; + + // 消息类型 + // - "TaskStarted": 任务已启动 + // - "SessionStarted": 会话已启动 + // - "TaskFailed": 任务失败 + // - "SessionFailed": 会话失败 + // - "SessionFinished": 会话结束 + string message_type = 4; + + // 状态码 + int32 status_code = 5; + + // 状态消息 (如 "OK") + string status_message = 6; + + // JSON 格式的识别结果 + // {"results": [...], "extra": {...}} + string result_json = 7; + + // 不知道是干啥的。。。 + int32 unknown_field_9 = 9; +} + +// 音频帧状态枚举 +enum FrameState { + FRAME_STATE_UNSPECIFIED = 0; + FRAME_STATE_FIRST = 1; + FRAME_STATE_MIDDLE = 3; + FRAME_STATE_LAST = 9; +} diff --git a/custom_components/hass_stt_doubao/doubaoime_asr/asr.py b/custom_components/hass_stt_doubao/doubaoime_asr/asr.py new file mode 100644 index 0000000..af11c57 --- /dev/null +++ b/custom_components/hass_stt_doubao/doubaoime_asr/asr.py @@ -0,0 +1,656 @@ +from __future__ import annotations + +import asyncio +import contextlib +from dataclasses import dataclass +from enum import Enum, auto +import json +from pathlib import Path +import ssl +import time +from typing import AsyncIterator, Callable, List, Optional, Union +import uuid +from pydantic import BaseModel, Field +import websockets +from websockets import ClientConnection + +from .config import ASRConfig, SessionConfig +from .audio import AudioEncoder +from .asr_pb2 import AsrRequest, AsrResponse as AsrResponsePb, FrameState + +# PCM 音频数据的类型别名 +AudioChunk = bytes + + +class ResponseType(Enum): + """ + ASR 响应类型 + """ + TASK_STARTED = auto() + SESSION_STARTED = auto() + SESSION_FINISHED = auto() + VAD_START = auto() + INTERIM_RESULT = auto() + FINAL_RESULT = auto() + HEARTBEAT = auto() + ERROR = auto() + UNKNOWN = auto() + + +@dataclass +class ASRResponse: + """ + ASR 响应 + """ + type: ResponseType + text: str = "" + is_final: bool = False + vad_start: bool = False + vad_finished: bool = False + packet_number: int = -1 + error_msg: str = "" + raw_json: Optional[dict] = None + + +class ASRError(Exception): + """ + ASR 错误 + """ + def __init__(self, message: str, response: Optional[ASRResponse] = None) -> None: + super().__init__(message) + self.response = response + + +class _SessionState(BaseModel): + """ + ASR 会话状态 + """ + request_id: str = Field(default_factory=lambda: str(uuid.uuid4())) + final_text: str = "" + is_finished: bool = False + error: Optional[ASRResponse] = None + + +class DoubaoASR: + """ + 豆包输入法 ASR 客户端 + """ + def __init__(self, config: Optional[ASRConfig] = None): + self.config = config + self._encoder = AudioEncoder(self.config) + self._ssl_context: Optional[ssl.SSLContext] = None + + async def __aenter__(self) -> DoubaoASR: + return self + + async def __aexit__(self, exc_type, exc_val, exc_tb) -> None: + pass + + async def _ensure_ssl_context(self): + """ + 确保 SSL 上下文已创建(在 executor 中运行避免阻塞) + """ + if self._ssl_context is None: + loop = asyncio.get_event_loop() + self._ssl_context = await loop.run_in_executor(None, ssl.create_default_context) + + async def transcribe(self, audio: Union[str, Path, bytes], *, realtime = False, on_interim: Callable[[str], None] = None) -> str: + """ + 非流式语音识别 + + :param audio: 音频文件路径或 PCM 字节数据 + :param on_interim: 可选的中间结果回调 + :return: 最终识别文本 + """ + final_text = "" + async for response in self.transcribe_stream(audio, realtime=realtime): + if response.type == ResponseType.INTERIM_RESULT and on_interim: + on_interim(response.text) + elif response.type == ResponseType.FINAL_RESULT: + final_text = response.text + elif response.type == ResponseType.ERROR: + raise ASRError(response.error_msg, response) + return final_text + + async def transcribe_stream(self, audio: Union[str, Path, bytes], *, realtime: bool = False) -> AsyncIterator[ASRResponse]: + """ + 流式语音识别(完整音频) + + :param audio: 音频文件路径或 PCM 字节数据 + :param realtime: 是否按实时速度发送 + :return: ASR 响应流,包括中间结果和最终结果 + """ + if isinstance(audio, (str, Path)): + pcm_data = self._encoder.convert_audio_to_pcm( + audio, self.config.sample_rate, self.config.channels, + ) + else: + pcm_data = audio + + opus_frames = await self._encoder.pcm_to_opus_frames(pcm_data) + state = _SessionState() + + # 预创建 SSL 上下文避免阻塞 + await self._ensure_ssl_context() + + # 获取 WebSocket URL + ws_url = await self.config.get_ws_url() + + try: + async with websockets.connect( + ws_url, + additional_headers=self.config.headers, + open_timeout=self.config.connect_timeout, + ssl=self._ssl_context, + ) as ws: + # 初始化会话 + async for resp in self._initialize_session(ws, state): + yield resp + + # 响应队列 + response_queue: asyncio.Queue[Optional[ASRResponse]] = asyncio.Queue() + + # 启动发送和接收任务 + send_task = asyncio.create_task( + self._send_audio(ws, opus_frames, state, realtime) + ) + recv_task = asyncio.create_task( + self._receive_responses(ws, state, response_queue) + ) + + try: + # 从队列中获取服务器响应 + while True: + try: + resp = await asyncio.wait_for( + response_queue.get(), + timeout=self.config.recv_timeout, + ) + if resp is None: # 结束标记 + break + + # 心跳包仅用于重置超时,不转发给用户 + if resp.type == ResponseType.HEARTBEAT: + continue + + yield resp + if resp.type == ResponseType.ERROR: + break + + except asyncio.TimeoutError: + break + + await send_task + finally: + send_task.cancel() + recv_task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await send_task + with contextlib.suppress(asyncio.CancelledError): + await recv_task + + except websockets.exceptions.WebSocketException as e: + raise ASRError(f"WebSocket 错误: {e}") from e + + async def transcribe_realtime( + self, + audio_source: AsyncIterator[AudioChunk], + ) -> AsyncIterator[ASRResponse]: + """ + 实时流式语音识别(支持麦克风等持续音频源) + + :param audio_source: PCM 音频数据的异步迭代器 + - 每个 chunk 应为 16-bit PCM 数据 + - 采样率和声道数应与 config 中配置一致 + - 迭代器结束时会自动发送 FinishSession + :return: ASR 响应流 + """ + state = _SessionState() + + # 预创建 SSL 上下文避免阻塞 + await self._ensure_ssl_context() + + # 获取 WebSocket URL + ws_url = await self.config.get_ws_url() + + try: + async with websockets.connect( + ws_url, + additional_headers=self.config.headers, + open_timeout=self.config.connect_timeout, + ssl=self._ssl_context, + ) as ws: + # 初始化会话 + async for resp in self._initialize_session(ws, state): + yield resp + + # 响应队列 + response_queue: asyncio.Queue[Optional[ASRResponse]] = asyncio.Queue() + + # 启动发送和接收任务 + send_task = asyncio.create_task( + self._send_audio_realtime(ws, audio_source, state) + ) + recv_task = asyncio.create_task( + self._receive_responses(ws, state, response_queue) + ) + + try: + # 实时模式不使用超时,依靠 WebSocket 层检测断开 + while True: + resp = await response_queue.get() + if resp is None: + break + + if resp.type == ResponseType.HEARTBEAT: + continue + + yield resp + if resp.type == ResponseType.ERROR: + break + + await send_task + finally: + send_task.cancel() + recv_task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await send_task + with contextlib.suppress(asyncio.CancelledError): + await recv_task + + except websockets.exceptions.WebSocketException as e: + raise ASRError(f"WebSocket 错误: {e}") from e + + async def _send_audio_realtime( + self, + ws: ClientConnection, + audio_source: AsyncIterator[AudioChunk], + state: _SessionState, + ): + """ + 从异步迭代器读取 PCM 数据并实时发送 + """ + # 预先获取编码器,避免在循环中触发阻塞导入 + encoder = await self._encoder.get_encoder() + + timestamp_ms = int(time.time() * 1000) + frame_index = 0 + pcm_buffer = b"" + + samples_per_frame = ( + self.config.sample_rate * self.config.frame_duration_ms // 1000 + ) + bytes_per_frame = samples_per_frame * 2 # 16-bit + + async for chunk in audio_source: + if state.is_finished: + break + + pcm_buffer += chunk + + # 当缓冲区有足够数据时,编码并发送 + while len(pcm_buffer) >= bytes_per_frame: + pcm_frame = pcm_buffer[:bytes_per_frame] + pcm_buffer = pcm_buffer[bytes_per_frame:] + + # 编码为 Opus + opus_frame = encoder.encode(pcm_frame, samples_per_frame) + + # 确定帧状态(实时模式下不知道最后一帧,使用 FIRST/MIDDLE) + if frame_index == 0: + frame_state = FrameState.FRAME_STATE_FIRST + else: + frame_state = FrameState.FRAME_STATE_MIDDLE + + msg = _build_asr_request( + opus_frame, + state.request_id, + frame_state, + timestamp_ms + frame_index * self.config.frame_duration_ms, + ) + await ws.send(msg) + frame_index += 1 + + # 迭代器结束,处理剩余数据 + if pcm_buffer and not state.is_finished: + # 补零到完整帧 + if len(pcm_buffer) < bytes_per_frame: + pcm_buffer += b"\x00" * (bytes_per_frame - len(pcm_buffer)) + + opus_frame = encoder.encode(pcm_buffer, samples_per_frame) + + msg = _build_asr_request( + opus_frame, + state.request_id, + FrameState.FRAME_STATE_LAST, + timestamp_ms + frame_index * self.config.frame_duration_ms, + ) + await ws.send(msg) + elif frame_index > 0 and not state.is_finished: + # 没有剩余数据,但需要发送一个 LAST 帧标记 + # 发送一个空的 LAST 帧(静音) + silent_frame = b"\x00" * bytes_per_frame + opus_frame = encoder.encode(silent_frame, samples_per_frame) + + msg = _build_asr_request( + opus_frame, + state.request_id, + FrameState.FRAME_STATE_LAST, + timestamp_ms + frame_index * self.config.frame_duration_ms, + ) + await ws.send(msg) + + # FinishSession + if not state.is_finished: + token = await self.config.get_token() + await ws.send(_build_finish_session(state.request_id, token)) + + async def _initialize_session(self, ws: ClientConnection, state: _SessionState) -> AsyncIterator[ASRResponse]: + """ + 初始化 ASR 会话 + """ + token = await self.config.get_token() + + # StartTask + await ws.send(_build_start_task(state.request_id, token)) + resp = await ws.recv() + parsed = _parse_response(resp) + if parsed.type == ResponseType.ERROR: + raise ASRError(f'StartTask 失败:{parsed.error_msg}', parsed) + yield parsed + + # StartSession + session_config = await self.config.get_session_config() + await ws.send( + _build_start_session(state.request_id, token, session_config) + ) + resp = await ws.recv() + parsed = _parse_response(resp) + if parsed.type == ResponseType.ERROR: + raise ASRError(f'StartSession 失败:{parsed.error_msg}', parsed) + yield parsed + + async def _send_audio( + self, + ws: ClientConnection, + opus_frames: List[bytes], + state: _SessionState, + realtime: bool, + ): + """ + 发送音频帧 + """ + timestamp_ms = int(time.time() * 1000) + frame_interval = self.config.frame_duration_ms / 1000.0 + + for i, opus_frame in enumerate(opus_frames): + if state.is_finished: + break + + if i == 0: + frame_state = FrameState.FRAME_STATE_FIRST + elif i == len(opus_frames) - 1: + frame_state = FrameState.FRAME_STATE_LAST + else: + frame_state = FrameState.FRAME_STATE_MIDDLE + + msg = _build_asr_request( + opus_frame, + state.request_id, + frame_state, + timestamp_ms + i * self.config.frame_duration_ms, + ) + await ws.send(msg) + + if realtime: + await asyncio.sleep(frame_interval) + + # FinishSession + token = await self.config.get_token() + await ws.send(_build_finish_session(state.request_id, token)) + + async def _receive_responses( + self, + ws: ClientConnection, + state: _SessionState, + queue: asyncio.Queue[Optional[ASRResponse]], + ): + """ + 接受响应并放入队列 + """ + try: + while not state.is_finished: + response = await ws.recv() + resp = _parse_response(response) + + if resp.type == ResponseType.ERROR: + state.error = resp + state.is_finished = True + await queue.put(resp) + break + elif resp.type == ResponseType.HEARTBEAT: + # 心跳包也放入队列,用于重置超时计时器 + await queue.put(resp) + elif resp.type == ResponseType.SESSION_FINISHED: + state.is_finished = True + await queue.put(resp) + break + elif resp.type == ResponseType.FINAL_RESULT: + state.final_text = resp.text + await queue.put(resp) + else: + await queue.put(resp) + + except websockets.exceptions.ConnectionClosed: + state.is_finished = True + finally: + # 结束信号 + await queue.put(None) + + + +def _build_start_task(request_id: str, token: str) -> bytes: + """构建 StartTask 消息 pb 数据""" + request = AsrRequest() + request.token = token + request.service_name = "ASR" + request.method_name = "StartTask" + request.request_id = request_id + return request.SerializeToString() + + +def _build_start_session(request_id: str, token: str, config: SessionConfig) -> bytes: + """构建 StartSession 消息 pb 数据""" + request = AsrRequest() + request.token = token + request.service_name = "ASR" + request.method_name = "StartSession" + request.request_id = request_id + request.payload = config.model_dump_json() + return request.SerializeToString() + + +def _build_finish_session(request_id: str, token: str) -> bytes: + """构建 FinishSession 消息 pb 数据""" + request = AsrRequest() + request.token = token + request.service_name = "ASR" + request.method_name = "FinishSession" + request.request_id = request_id + return request.SerializeToString() + + +def _build_asr_request( + audio_data: bytes, + request_id: str, + frame_state: FrameState, + timestamp_ms: int, +) -> bytes: + request = AsrRequest() + metadata = json.dumps({"extra": {}, "timestamp_ms": timestamp_ms}) + + request.service_name = "ASR" + request.method_name = "TaskRequest" + request.payload = metadata + request.audio_data = audio_data + request.request_id = request_id + request.frame_state = frame_state + return request.SerializeToString() + + + +def _parse_response(data: bytes) -> ASRResponse: + """解析 ASR 响应 (使用 protobuf)""" + pb = AsrResponsePb() + pb.ParseFromString(data) + + message_type = pb.message_type + result_json = pb.result_json # 字段 7: 识别结果 JSON + status_message = pb.status_message # 字段 6: 状态消息 + + # 根据 message_type 判断响应类型 + if message_type == "TaskStarted": + return ASRResponse(type=ResponseType.TASK_STARTED) + + if message_type == "SessionStarted": + return ASRResponse(type=ResponseType.SESSION_STARTED) + + if message_type == "SessionFinished": + return ASRResponse(type=ResponseType.SESSION_FINISHED) + + if message_type in ("TaskFailed", "SessionFailed"): + return ASRResponse(type=ResponseType.ERROR, error_msg=status_message) + + # 识别结果在 result_json 字段(字段 7) + if not result_json: + return ASRResponse(type=ResponseType.UNKNOWN) + + try: + json_data = json.loads(result_json) + except json.JSONDecodeError: + return ASRResponse(type=ResponseType.UNKNOWN) + + results = json_data.get("results") + extra = json_data.get("extra", {}) + + # 无 results,可能是心跳包 + if results is None: + return ASRResponse( + type=ResponseType.HEARTBEAT, + packet_number=extra.get("packet_number", -1), + raw_json=json_data, + ) + + # VAD 开始 + if extra.get("vad_start"): + return ASRResponse(type=ResponseType.VAD_START, vad_start=True, raw_json=json_data) + + # 解析识别结果 + text = "" + is_interim = True + vad_finished = False + nonstream_result = False + + for r in results: + if r.get("text"): + text = r.get("text") + if r.get("is_interim") is False: + is_interim = False + if r.get("is_vad_finished"): + vad_finished = True + if r.get("extra", {}).get("nonstream_result"): + nonstream_result = True + + # 最终结果 + if nonstream_result or (not is_interim and vad_finished): + return ASRResponse( + type=ResponseType.FINAL_RESULT, + text=text, + is_final=True, + vad_finished=vad_finished, + raw_json=json_data, + ) + + # 中间结果 + return ASRResponse( + type=ResponseType.INTERIM_RESULT, + text=text, + is_final=False, + raw_json=json_data, + ) + + +# ============= +# 便捷函数 +# ============= + + +async def transcribe( + audio: str | Path | bytes, + *, + config: ASRConfig | None = None, + on_interim: Callable[[str], None] | None = None, + realtime: bool = False, +) -> str: + """ + 便捷函数:非流式语音识别 + + Args: + audio: 音频文件路径或 PCM 字节数据 + config: ASR 配置(可选) + on_interim: 中间结果回调(可选) + realtime: 是否模拟实时语音输入 + - True: 按音频实际时长发送,每帧间插入延迟,模拟实时的语音输入 + - False(默认): 尽快发送所有帧,会更快拿到结果(不知道会不会被风控) + + Returns: + 最终识别文本 + """ + async with DoubaoASR(config) as asr: + return await asr.transcribe(audio, on_interim=on_interim, realtime=realtime) + + +async def transcribe_stream( + audio: str | Path | bytes, + *, + config: ASRConfig | None = None, + realtime: bool = False, +) -> AsyncIterator[ASRResponse]: + """ + 便捷函数:流式语音识别(完整音频) + + Args: + audio: 音频文件路径或 PCM 字节数据 + config: ASR 配置(可选) + realtime: 是否模拟实时语音输入 + - True: 按音频实际时长发送,每帧间插入延迟,模拟实时的语音输入 + - False(默认): 尽快发送所有帧,会更快拿到结果(不知道会不会被风控) + + Yields: + ASRResponse 对象 + """ + async with DoubaoASR(config) as asr: + async for resp in asr.transcribe_stream(audio, realtime=realtime): + yield resp + + +async def transcribe_realtime( + audio_source: AsyncIterator[AudioChunk], + *, + config: ASRConfig | None = None, +) -> AsyncIterator[ASRResponse]: + """ + 便捷函数:实时流式语音识别(支持麦克风等持续音频源) + + Args: + audio_source: PCM 音频数据的异步迭代器 + - 每个 chunk 应为 16-bit PCM 数据 + - 采样率和声道数应与 config 中配置一致 + config: ASR 配置(可选) + + Yields: + ASRResponse 对象 + """ + async with DoubaoASR(config) as asr: + async for resp in asr.transcribe_realtime(audio_source): + yield resp diff --git a/custom_components/hass_stt_doubao/doubaoime_asr/asr_pb2.py b/custom_components/hass_stt_doubao/doubaoime_asr/asr_pb2.py new file mode 100644 index 0000000..e98dfa2 --- /dev/null +++ b/custom_components/hass_stt_doubao/doubaoime_asr/asr_pb2.py @@ -0,0 +1,40 @@ +# -*- coding: utf-8 -*- +# Generated by the protocol buffer compiler. DO NOT EDIT! +# NO CHECKED-IN PROTOBUF GENCODE +# source: asr.proto +# Protobuf Python Version: 6.31.1 +"""Generated protocol buffer code.""" +from google.protobuf import descriptor as _descriptor +from google.protobuf import descriptor_pool as _descriptor_pool +from google.protobuf import runtime_version as _runtime_version +from google.protobuf import symbol_database as _symbol_database +from google.protobuf.internal import builder as _builder +_runtime_version.ValidateProtobufRuntimeVersion( + _runtime_version.Domain.PUBLIC, + 6, + 31, + 1, + '', + 'asr.proto' +) +# @@protoc_insertion_point(imports) + +_sym_db = _symbol_database.Default() + + + + +DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\tasr.proto\x12\x03\x61sr\"\xa5\x01\n\nAsrRequest\x12\r\n\x05token\x18\x02 \x01(\t\x12\x14\n\x0cservice_name\x18\x03 \x01(\t\x12\x13\n\x0bmethod_name\x18\x05 \x01(\t\x12\x0f\n\x07payload\x18\x06 \x01(\t\x12\x12\n\naudio_data\x18\x07 \x01(\x0c\x12\x12\n\nrequest_id\x18\x08 \x01(\t\x12$\n\x0b\x66rame_state\x18\t \x01(\x0e\x32\x0f.asr.FrameState\"\xb9\x01\n\x0b\x41srResponse\x12\x12\n\nrequest_id\x18\x01 \x01(\t\x12\x0f\n\x07task_id\x18\x02 \x01(\t\x12\x14\n\x0cservice_name\x18\x03 \x01(\t\x12\x14\n\x0cmessage_type\x18\x04 \x01(\t\x12\x13\n\x0bstatus_code\x18\x05 \x01(\x05\x12\x16\n\x0estatus_message\x18\x06 \x01(\t\x12\x13\n\x0bresult_json\x18\x07 \x01(\t\x12\x17\n\x0funknown_field_9\x18\t \x01(\x05*n\n\nFrameState\x12\x1b\n\x17\x46RAME_STATE_UNSPECIFIED\x10\x00\x12\x15\n\x11\x46RAME_STATE_FIRST\x10\x01\x12\x16\n\x12\x46RAME_STATE_MIDDLE\x10\x03\x12\x14\n\x10\x46RAME_STATE_LAST\x10\tb\x06proto3') + +_globals = globals() +_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) +_builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, 'asr_pb2', _globals) +if not _descriptor._USE_C_DESCRIPTORS: + DESCRIPTOR._loaded_options = None + _globals['_FRAMESTATE']._serialized_start=374 + _globals['_FRAMESTATE']._serialized_end=484 + _globals['_ASRREQUEST']._serialized_start=19 + _globals['_ASRREQUEST']._serialized_end=184 + _globals['_ASRRESPONSE']._serialized_start=187 + _globals['_ASRRESPONSE']._serialized_end=372 +# @@protoc_insertion_point(module_scope) diff --git a/custom_components/hass_stt_doubao/doubaoime_asr/asr_pb2.pyi b/custom_components/hass_stt_doubao/doubaoime_asr/asr_pb2.pyi new file mode 100644 index 0000000..d988272 --- /dev/null +++ b/custom_components/hass_stt_doubao/doubaoime_asr/asr_pb2.pyi @@ -0,0 +1,55 @@ +from google.protobuf.internal import enum_type_wrapper as _enum_type_wrapper +from google.protobuf import descriptor as _descriptor +from google.protobuf import message as _message +from typing import ClassVar as _ClassVar, Optional as _Optional, Union as _Union + +DESCRIPTOR: _descriptor.FileDescriptor + +class FrameState(int, metaclass=_enum_type_wrapper.EnumTypeWrapper): + __slots__ = () + FRAME_STATE_UNSPECIFIED: _ClassVar[FrameState] + FRAME_STATE_FIRST: _ClassVar[FrameState] + FRAME_STATE_MIDDLE: _ClassVar[FrameState] + FRAME_STATE_LAST: _ClassVar[FrameState] +FRAME_STATE_UNSPECIFIED: FrameState +FRAME_STATE_FIRST: FrameState +FRAME_STATE_MIDDLE: FrameState +FRAME_STATE_LAST: FrameState + +class AsrRequest(_message.Message): + __slots__ = ("token", "service_name", "method_name", "payload", "audio_data", "request_id", "frame_state") + TOKEN_FIELD_NUMBER: _ClassVar[int] + SERVICE_NAME_FIELD_NUMBER: _ClassVar[int] + METHOD_NAME_FIELD_NUMBER: _ClassVar[int] + PAYLOAD_FIELD_NUMBER: _ClassVar[int] + AUDIO_DATA_FIELD_NUMBER: _ClassVar[int] + REQUEST_ID_FIELD_NUMBER: _ClassVar[int] + FRAME_STATE_FIELD_NUMBER: _ClassVar[int] + token: str + service_name: str + method_name: str + payload: str + audio_data: bytes + request_id: str + frame_state: FrameState + def __init__(self, token: _Optional[str] = ..., service_name: _Optional[str] = ..., method_name: _Optional[str] = ..., payload: _Optional[str] = ..., audio_data: _Optional[bytes] = ..., request_id: _Optional[str] = ..., frame_state: _Optional[_Union[FrameState, str]] = ...) -> None: ... + +class AsrResponse(_message.Message): + __slots__ = ("request_id", "task_id", "service_name", "message_type", "status_code", "status_message", "result_json", "unknown_field_9") + REQUEST_ID_FIELD_NUMBER: _ClassVar[int] + TASK_ID_FIELD_NUMBER: _ClassVar[int] + SERVICE_NAME_FIELD_NUMBER: _ClassVar[int] + MESSAGE_TYPE_FIELD_NUMBER: _ClassVar[int] + STATUS_CODE_FIELD_NUMBER: _ClassVar[int] + STATUS_MESSAGE_FIELD_NUMBER: _ClassVar[int] + RESULT_JSON_FIELD_NUMBER: _ClassVar[int] + UNKNOWN_FIELD_9_FIELD_NUMBER: _ClassVar[int] + request_id: str + task_id: str + service_name: str + message_type: str + status_code: int + status_message: str + result_json: str + unknown_field_9: int + def __init__(self, request_id: _Optional[str] = ..., task_id: _Optional[str] = ..., service_name: _Optional[str] = ..., message_type: _Optional[str] = ..., status_code: _Optional[int] = ..., status_message: _Optional[str] = ..., result_json: _Optional[str] = ..., unknown_field_9: _Optional[int] = ...) -> None: ... diff --git a/custom_components/hass_stt_doubao/doubaoime_asr/audio.py b/custom_components/hass_stt_doubao/doubaoime_asr/audio.py new file mode 100644 index 0000000..d6650f5 --- /dev/null +++ b/custom_components/hass_stt_doubao/doubaoime_asr/audio.py @@ -0,0 +1,77 @@ +from typing import Optional, List, Union, TYPE_CHECKING +from pathlib import Path +import asyncio + +import miniaudio + +from .config import ASRConfig + +if TYPE_CHECKING: + import opuslib + + +class AudioEncoder: + """ + 进行音频格式转换 + """ + def __init__(self, config: ASRConfig) -> None: + self.config = config + self._encoder: Optional["opuslib.Encoder"] = None + self._opuslib_module = None + + async def _ensure_opuslib(self): + """在 executor 中导入 opuslib,避免阻塞事件循环""" + if self._opuslib_module is None: + loop = asyncio.get_event_loop() + self._opuslib_module = await loop.run_in_executor(None, self._import_opuslib) + + @staticmethod + def _import_opuslib(): + """在独立线程中导入 opuslib""" + import opuslib + return opuslib + + async def get_encoder(self) -> "opuslib.Encoder": + """异步获取编码器""" + if self._encoder is None: + await self._ensure_opuslib() + self._encoder = self._opuslib_module.Encoder( + self.config.sample_rate, + self.config.channels, + self._opuslib_module.APPLICATION_AUDIO, + ) + return self._encoder + + async def pcm_to_opus_frames(self, pcm_data: bytes) -> List[bytes]: + """将 PCM 数据转换为 Opus 帧(异步)""" + encoder = await self.get_encoder() + + samples_per_frame = ( + self.config.sample_rate * self.config.frame_duration_ms // 1000 + ) + bytes_per_frame = samples_per_frame * 2 # 16-bit + + frames = [] + for i in range(0, len(pcm_data), bytes_per_frame): + chunk = pcm_data[i : i + bytes_per_frame] + if len(chunk) < bytes_per_frame: + chunk = chunk + b"\x00" * (bytes_per_frame - len(chunk)) + + opus_frame = encoder.encode(chunk, samples_per_frame) + frames.append(opus_frame) + + return frames + + @staticmethod + def convert_audio_to_pcm( + audio_path: Union[Path, str], + sample_rate: int = 16000, + channels: int = 1, + ) -> bytes: + decoded = miniaudio.decode_file( + str(audio_path), + output_format=miniaudio.SampleFormat.SIGNED16, + nchannels=channels, + sample_rate=sample_rate, + ) + return decoded.samples.tobytes() diff --git a/custom_components/hass_stt_doubao/doubaoime_asr/config.py b/custom_components/hass_stt_doubao/doubaoime_asr/config.py new file mode 100644 index 0000000..2134d53 --- /dev/null +++ b/custom_components/hass_stt_doubao/doubaoime_asr/config.py @@ -0,0 +1,231 @@ +from dataclasses import dataclass, field +import asyncio +import json +from pathlib import Path +from typing import Optional, Union +from pydantic import BaseModel +import aiofiles + +from .constants import WEBSOCKET_URL, USER_AGENT, AID +from .device import DeviceCredentials, register_device, get_asr_token + + +class _AudioInfo(BaseModel): + channel: int + format: str + sample_rate: int + + +class _SessionExtraConfig(BaseModel): + app_name: str + cell_compress_rate: int + did: str + enable_asr_threepass: bool + enable_asr_twopass: bool + input_mode: str + + +class SessionConfig(BaseModel): + """ + ASR 任务开始前需要初始化 Session 的配置 + """ + audio_info: _AudioInfo + enable_punctuation: bool + enable_speech_rejection: bool + extra: _SessionExtraConfig + + +@dataclass +class ASRConfig: + """ + ASR 配置 + + 如果不提供 device_id 和 token,将自动注册设备并获取 token。 + + 示例: + # 自动获取凭据(首次使用时会注册设备,不持久化) + config = ASRConfig() + + # 使用已有凭据 + config = ASRConfig(device_id="1234567890123456", token="MyToken123") + + # 使用凭据文件(推荐,首次注册后自动缓存) + config = ASRConfig(credential_path="~/.config/doubao-asr/credentials.json") + + # 凭据文件 + 覆盖部分参数 + config = ASRConfig( + credential_path="~/.config/doubao-asr/credentials.json", + token="NewToken", # 覆盖文件中的 token + ) + """ + url: str = WEBSOCKET_URL + aid: str = AID + user_agent: str = USER_AGENT + + device_id: Optional[str] = None # 空则自动获取 + token: Optional[str] = None # 空则自动获取 + credential_path: Union[str, Path, None] = None + """ + 凭据文件路径 + """ + + # 这些都是客户端给的默认值,挺通用的。其实我也没尝试过改了服务器会不会认 + # 音频配置 + sample_rate: int = 16000 + channels: int = 1 + frame_duration_ms: int = 20 + + # 会话配置 + enable_punctuation: bool = True + enable_speech_rejection: bool = False + enable_asr_twopass: bool = True + enable_asr_threepass: bool = True + # 这里是输入法当前作用在哪个应用上 + # 可能服务器会根据当前所使用的程序调整不同的语音识别策略?? + # 这里用 Chrome 浏览器,算是比较通用的了吧? + app_name: str = "com.android.chrome" + + # 连接配置 + connect_timeout: float = 10.0 + recv_timeout: float = 10.0 + + # 内部状态 + _credentials: Optional[DeviceCredentials] = field(default=None, repr=None) + _initialized: bool = field(default=False, repr=False) + + async def _load_credentials_from_file(self) -> Optional[DeviceCredentials]: + """ + 从缓存文件中加载凭据信息(异步) + """ + if self.credential_path is None: + return None + + path = Path(self.credential_path).expanduser() + if not path.exists(): + return None + + try: + async with aiofiles.open(path, 'r', encoding='utf-8') as f: + content = await f.read() + data = json.loads(content) + return DeviceCredentials(**data) + + except (json.JSONDecodeError, OSError): + return None + + async def _save_credentials_to_file(self, creds: DeviceCredentials): + """ + 保存凭据至缓存文件(异步) + """ + if self.credential_path is None: + return + + path = Path(self.credential_path).expanduser() + + # 使用 run_in_executor 来执行同步的 mkdir 操作 + loop = asyncio.get_event_loop() + await loop.run_in_executor(None, lambda: path.parent.mkdir(parents=True, exist_ok=True)) + + async with aiofiles.open(path, 'w', encoding='utf-8') as f: + content = json.dumps(creds.model_dump(), indent=2, ensure_ascii=False) + await f.write(content) + + async def async_ensure_credentials(self): + """ + 确保凭据已初始化(异步) + + 优先级: + 1. 直接传入的 device_id/token 参数(最高优先级) + 2. credential_path 文件中的值 + 3. 自动注册获取(最低优先级) + + 如果指定了 credential_path 且文件不存在,会注册设备并保存到文件。 + """ + if self._initialized: + return + + # 保存直接通过参数传入的凭据,用于进行覆盖 + user_device_id = self.device_id + user_token = self.token + + # 尝试从文件中加载凭据 + file_creds = await self._load_credentials_from_file() + if file_creds: + self._credentials = file_creds + # 使用文件中的值作为默认 + if self.device_id is None: + self.device_id = file_creds.device_id + if self.token is None: + self.token = file_creds.token + + # 如果 device_id 仍为 None, 则注册设备 + need_save = False + if self.device_id is None: + # 在 executor 中运行同步的 register_device + loop = asyncio.get_event_loop() + self._credentials = await loop.run_in_executor(None, register_device) + self.device_id = self._credentials.device_id + need_save = True + + # 如果 token 仍为 None, 则获取 token + if self.token is None: + cdid = self._credentials.cdid if self._credentials else None + # 在 executor 中运行同步的 get_asr_token + loop = asyncio.get_event_loop() + self.token = await loop.run_in_executor(None, get_asr_token, self.device_id, cdid) + + # 如果指定了 credential_path 且有新注册的凭据,则保存至文件 + if self.credential_path and need_save and self._credentials: + self._credentials.token = self.token + await self._save_credentials_to_file(self._credentials) + + # 覆盖用户传入的参数 + if user_device_id is not None: + self.device_id = user_device_id + + if user_token is not None: + self.token = user_token + + self._initialized = True + + async def get_ws_url(self) -> str: + """获取 WebSocket URL(异步)""" + await self.async_ensure_credentials() + return f'{self.url}?aid={self.aid}&device_id={self.device_id}' + + @property + def headers(self) -> dict[str, str]: + return { + "User-Agent": self.user_agent, + "proto-version": "v2", + "x-custom-keepalive": "true" + } + + async def get_session_config(self) -> SessionConfig: + """获取会话配置(异步)""" + await self.async_ensure_credentials() + audio_info = _AudioInfo( + channel=self.channels, + format="speech_opus", + sample_rate=self.sample_rate, + ) + extra = _SessionExtraConfig( + app_name=self.app_name, + cell_compress_rate=8, + did=self.device_id, + enable_asr_threepass=self.enable_asr_threepass, + enable_asr_twopass=self.enable_asr_twopass, + input_mode="tool", + ) + + return SessionConfig( + audio_info=audio_info, + enable_punctuation=self.enable_punctuation, + enable_speech_rejection=self.enable_speech_rejection, + extra=extra, + ) + + async def get_token(self) -> str: + """获取 token(异步)""" + await self.async_ensure_credentials() + return self.token \ No newline at end of file diff --git a/custom_components/hass_stt_doubao/doubaoime_asr/constants.py b/custom_components/hass_stt_doubao/doubaoime_asr/constants.py new file mode 100644 index 0000000..5dea4ba --- /dev/null +++ b/custom_components/hass_stt_doubao/doubaoime_asr/constants.py @@ -0,0 +1,43 @@ +# 设备注册 API URL +REGISTER_URL = "https://log.snssdk.com/service/2/device_register/" + +# Settings API URL (获取 Token) +SETTINGS_URL = "https://is.snssdk.com/service/settings/v3/" + +# ASR WebSocket URL +WEBSOCKET_URL = "wss://frontier-audio-ime-ws.doubao.com/ocean/api/v1/ws" + +# 豆包输入法的 APP ID +AID = 401734 + +# 应用配置 (豆包输入法) +APP_CONFIG = { + "aid": AID, + "app_name": "oime", + "version_code": 100102018, + "version_name": "1.1.2", + "manifest_version_code": 100102018, + "update_version_code": 100102018, + "channel": "official", + "package": "com.bytedance.android.doubaoime", +} + +# 默认设备配置 (模拟 Pixel 7 Pro) +DEFAULT_DEVICE_CONFIG = { + "device_platform": "android", + "os": "android", + "os_api": "34", + "os_version": "16", + "device_type": "Pixel 7 Pro", + "device_brand": "google", + "device_model": "Pixel 7 Pro", + "resolution": "1080*2400", + "dpi": "420", + "language": "zh", + "timezone": 8, + "access": "wifi", + "rom": "UP1A.231005.007", + "rom_version": "UP1A.231005.007", +} + +USER_AGENT = "com.bytedance.android.doubaoime/100102018 (Linux; U; Android 16; en_US; Pixel 7 Pro; Build/BP2A.250605.031.A2; Cronet/TTNetVersion:94cf429a 2025-11-17 QuicVersion:1f89f732 2025-05-08)" \ No newline at end of file diff --git a/custom_components/hass_stt_doubao/doubaoime_asr/device.py b/custom_components/hass_stt_doubao/doubaoime_asr/device.py new file mode 100644 index 0000000..ac2446f --- /dev/null +++ b/custom_components/hass_stt_doubao/doubaoime_asr/device.py @@ -0,0 +1,321 @@ +""" +设备初始化相关 + +需要先根据客户端配置在豆包服务器注册设备,获取 device_id, install_id, token 等信息 +""" +import hashlib +from pydantic import BaseModel, ConfigDict, Field +from typing import Optional + +import secrets +import requests +import time +import uuid + +from .constants import APP_CONFIG, DEFAULT_DEVICE_CONFIG, USER_AGENT, REGISTER_URL, SETTINGS_URL + + +class DeviceCredentials(BaseModel): + """ + 设备凭据,用于缓存设备信息 + 主要是 `device_id` 和 `token`. 其他貌似不影响 + """ + device_id: Optional[str] = None + install_id: Optional[str] = None + cdid: Optional[str] = None + openudid: Optional[str] = None + clientudid: Optional[str] = None + token: Optional[str] = "" + """ + 用于 ASR 的 token + """ + + +class DeviceRegisterHeaderField(BaseModel): + """ + 设备注册接口用到的请求体 header 字段 + """ + # 设备标识(注册时为 0,注册后更新) + device_id: int = 0 + install_id: int = 0 + + # 应用配置 + aid: int + """app id,固定值""" + app_name: str + version_code: int + version_name: str + manifest_version_code: int + update_version_code: int + channel: str + package: str + + # 设备平台信息 + device_platform: str + os: str + os_api: str + os_version: str + device_type: str + device_brand: str + device_model: str + resolution: str + dpi: str + language: str + timezone: int + access: str + rom: str + rom_version: str + + # 设备唯一标识 + openudid: str + clientudid: str + cdid: str + + # 地区与时区 + region: str = "CN" + tz_name: str = "Asia/Shanghai" + tz_offset: int = 28800 + sim_region: str = "cn" + carrier_region: str = "cn" + + # 其他设备信息 + cpu_abi: str = "arm64-v8a" + build_serial: str = "unknown" + not_request_sender: int = 0 + sig_hash: str = "" + google_aid: str = "" + mc: str = "" + serial_number: str = "" + + @classmethod + def default(cls, cdid: Optional[str] = None, openudid: Optional[str] = None, clientudid: Optional[str] = None) -> "DeviceRegisterHeaderField": + """ + 使用默认配置构建设备注册 Header + """ + + return cls( + **APP_CONFIG, + **DEFAULT_DEVICE_CONFIG, + cdid=cdid or _generate_cdid(), + openudid=openudid or _generate_openudid(), + clientudid=clientudid or _generate_clientudid(), + ) + + +class DeviceRegisterBody(BaseModel): + """ + 设备注册接口用到的完整的请求体 + """ + model_config = ConfigDict(populate_by_name=True, serialize_by_alias=True) + + magic_tag: str = "ss_app_log" + header: DeviceRegisterHeaderField + gen_time: int = Field(default_factory=lambda: int(time.time() * 1000), serialization_alias="_gen_time") + + @classmethod + def new(cls, header: DeviceRegisterHeaderField): + return cls(header=header) + + +class DeviceRegisterParams(BaseModel): + """ + 设备注册接口的 URL Params + """ + model_config = ConfigDict(populate_by_name=True, serialize_by_alias=True) + + device_platform: str + os: str + ssmix: str = "a" + rticket: int = Field(default_factory=lambda: int(time.time() * 1000), serialization_alias="_rticket") + cdid: str + + # 应用配置 + channel: str + aid: str + app_name: str + version_code: str + version_name: str + manifest_version_code: str + update_version_code: str + + # 设备信息 + resolution: str + dpi: str + device_type: str + device_brand: str + language: str + os_api: str + os_version: str + ac: str = "wifi" + + @classmethod + def default(cls, cdid: str) -> "DeviceRegisterParams": + """ + 使用默认配置构建 URL Params + """ + + app_config = { + **{k: APP_CONFIG[k] for k in ("channel", "app_name", "version_name")}, + "aid": str(APP_CONFIG["aid"]), + "version_code": str(APP_CONFIG["version_code"]), + "manifest_version_code": str(APP_CONFIG["manifest_version_code"]), + "update_version_code": str(APP_CONFIG["update_version_code"]), + } + + device_keys = ("device_platform", "os", "resolution", "dpi", "device_type", + "device_brand", "language", "os_api", "os_version") + device_config = {k: DEFAULT_DEVICE_CONFIG[k] for k in device_keys} + + return cls(cdid=cdid, **app_config, **device_config) + + +class DeviceRegisterResponse(BaseModel): + server_time: int + device_id: int + install_id: int + new_user: Optional[int] = None + device_id_str: Optional[str] = None + install_id_str: Optional[str] = None + ssid: Optional[str] = None + device_token: Optional[str] = None + + +class SettingsParams(BaseModel): + """ + Settings API 的 URL Params(用于获取 ASR token) + """ + model_config = ConfigDict(populate_by_name=True, serialize_by_alias=True) + + device_platform: str = "android" + os: str = "android" + ssmix: str = "a" + rticket: str = Field(default_factory=lambda: str(int(time.time() * 1000)), serialization_alias="_rticket") + cdid: str + channel: str + aid: str + app_name: str + version_code: str + version_name: str + device_id: str + + @classmethod + def default(cls, device_id: str, cdid: str) -> "SettingsParams": + """ + 使用默认配置构建 Settings Params + """ + return cls( + cdid=cdid, + device_id=device_id, + channel=APP_CONFIG["channel"], + aid=str(APP_CONFIG["aid"]), + app_name=APP_CONFIG["app_name"], + version_code=str(APP_CONFIG["version_code"]), + version_name=APP_CONFIG["version_name"], + ) + + +class _AsrConfig(BaseModel): + """ASR 配置""" + app_key: str + + +class _Settings(BaseModel): + """Settings 配置""" + asr_config: _AsrConfig + + +class _SettingsData(BaseModel): + """Settings 数据""" + settings: _Settings + + +class SettingsResponse(BaseModel): + """Settings API 响应""" + data: _SettingsData + message: str + + @property + def app_key(self) -> str: + """获取 ASR app_key (token)""" + return self.data.settings.asr_config.app_key + + +def _generate_openudid() -> str: + return secrets.token_hex(8) + + +def _generate_cdid() -> str: + return str(uuid.uuid4()) + + +def _generate_clientudid() -> str: + return str(uuid.uuid4()) + + +def register_device() -> DeviceCredentials: + """ + 首次使用,注册设备获取 device_id + """ + cdid = _generate_cdid() + openudid = _generate_openudid() + clientudid = _generate_clientudid() + + header = DeviceRegisterHeaderField.default(cdid=cdid, openudid=openudid, clientudid=clientudid) + body = DeviceRegisterBody.new(header) + params = DeviceRegisterParams.default(cdid) + + headers = { + "User-Agent": USER_AGENT, + } + + response = requests.post( + REGISTER_URL, + params=params.model_dump(), + json=body.model_dump(), + headers=headers, + verify=False, # 禁用SSL证书验证以兼容某些环境 + ) + + response.raise_for_status() + response_json = response.json() + response_data = DeviceRegisterResponse(**response_json) + + if response_data.device_id and response_data.device_id != 0: + return DeviceCredentials( + device_id=str(response_data.device_id), + install_id=str(response_data.install_id), + cdid=cdid, + openudid=openudid, + clientudid=clientudid, + ) + + +def get_asr_token(device_id: str, cdid: str) -> str: + """ + 获取 ASR 请求所需的 token + """ + if cdid is None: + cdid = _generate_cdid() + + params = SettingsParams.default(device_id, cdid) + body_str = "body=null" + x_ss_stub = hashlib.md5(body_str.encode()).hexdigest().upper() + + headers = { + "User-Agent": USER_AGENT, + "x-ss-stub": x_ss_stub, + } + + response = requests.post( + SETTINGS_URL, + params=params, + data=body_str, + headers=headers, + verify=False, # 禁用SSL证书验证以兼容某些环境 + ) + + response.raise_for_status() + response_json = response.json() + response_data = SettingsResponse(**response_json) + + return response_data.app_key diff --git a/custom_components/hass_stt_doubao/manifest.json b/custom_components/hass_stt_doubao/manifest.json new file mode 100644 index 0000000..37dd8ea --- /dev/null +++ b/custom_components/hass_stt_doubao/manifest.json @@ -0,0 +1,18 @@ +{ + "domain": "hass_stt_doubao", + "name": "Doubao Speech to Text", + "version": "1.0.0", + "documentation": "https://github.com/yourusername/hass_stt_doubao", + "issue_tracker": "https://github.com/yourusername/hass_stt_doubao/issues", + "dependencies": [], + "codeowners": ["@xyzmos"], + "requirements": [ + "miniaudio>=1.61", + "opuslib>=3.0.1", + "protobuf>=3.20.0", + "pydantic>=2.0.0", + "websockets>=12.0" + ], + "config_flow": true, + "iot_class": "cloud_polling" +} diff --git a/custom_components/hass_stt_doubao/requirements.txt b/custom_components/hass_stt_doubao/requirements.txt new file mode 100644 index 0000000..e4dd097 --- /dev/null +++ b/custom_components/hass_stt_doubao/requirements.txt @@ -0,0 +1,6 @@ +aiofiles>=23.0.0 +miniaudio>=1.61 +opuslib>=3.0.1 +protobuf>=3.20.0 +pydantic>=2.0.0 +websockets>=12.0 diff --git a/custom_components/hass_stt_doubao/strings.json b/custom_components/hass_stt_doubao/strings.json new file mode 100644 index 0000000..b0b6917 --- /dev/null +++ b/custom_components/hass_stt_doubao/strings.json @@ -0,0 +1,33 @@ +{ + "config": { + "step": { + "user": { + "title": "配置 Doubao 语音识别", + "description": "设置 Doubao STT 集成", + "data": { + "credential_path": "凭据文件路径", + "enable_punctuation": "启用标点符号" + } + } + }, + "error": { + "cannot_connect": "无法连接到 Doubao 服务,请检查网络连接或凭据配置", + "unknown": "发生未知错误" + }, + "abort": { + "already_configured": "该设备已经配置过了" + } + }, + "options": { + "step": { + "init": { + "title": "Doubao STT 选项", + "description": "修改 Doubao STT 配置选项", + "data": { + "credential_path": "凭据文件路径", + "enable_punctuation": "启用标点符号" + } + } + } + } +} diff --git a/custom_components/hass_stt_doubao/stt.py b/custom_components/hass_stt_doubao/stt.py new file mode 100644 index 0000000..23a8f18 --- /dev/null +++ b/custom_components/hass_stt_doubao/stt.py @@ -0,0 +1,162 @@ +"""Support for Doubao Speech-to-Text service.""" +from __future__ import annotations + +import asyncio +import logging +from typing import AsyncIterable + +from homeassistant.components.stt import ( + AudioBitRates, + AudioChannels, + AudioCodecs, + AudioFormats, + AudioSampleRates, + SpeechMetadata, + SpeechResult, + SpeechResultState, + SpeechToTextEntity, +) +from homeassistant.config_entries import ConfigEntry +from homeassistant.core import HomeAssistant +from homeassistant.helpers.entity_platform import AddEntitiesCallback + +from .const import ( + DOMAIN, + CONF_CREDENTIAL_PATH, + CONF_ENABLE_PUNCTUATION, + SUPPORTED_LANGUAGES, +) +from .doubaoime_asr import ASRConfig, ASRError, DoubaoASR, ResponseType + +_LOGGER = logging.getLogger(__name__) + + +async def async_setup_entry( + hass: HomeAssistant, + config_entry: ConfigEntry, + async_add_entities: AddEntitiesCallback, +) -> None: + """Set up Doubao STT from a config entry.""" + config_data = hass.data[DOMAIN][config_entry.entry_id] + + async_add_entities( + [DoubaoSTTEntity(config_entry, config_data)], + True, + ) + + +class DoubaoSTTEntity(SpeechToTextEntity): + """Doubao Speech-to-Text entity.""" + + def __init__( + self, + config_entry: ConfigEntry, + config_data: dict, + ) -> None: + """Initialize Doubao STT entity.""" + self._config_entry = config_entry + self._credential_path = config_data[CONF_CREDENTIAL_PATH] + self._enable_punctuation = config_data[CONF_ENABLE_PUNCTUATION] + self._attr_name = "Doubao STT" + self._attr_unique_id = f"{config_entry.entry_id}_stt" + + @property + def supported_languages(self) -> list[str]: + """Return a list of supported languages.""" + return SUPPORTED_LANGUAGES + + @property + def supported_formats(self) -> list[AudioFormats]: + """Return a list of supported formats.""" + return [AudioFormats.WAV] + + @property + def supported_codecs(self) -> list[AudioCodecs]: + """Return a list of supported codecs.""" + return [AudioCodecs.PCM] + + @property + def supported_bit_rates(self) -> list[AudioBitRates]: + """Return a list of supported bit rates.""" + return [AudioBitRates.BITRATE_16] + + @property + def supported_sample_rates(self) -> list[AudioSampleRates]: + """Return a list of supported sample rates.""" + return [AudioSampleRates.SAMPLERATE_16000] + + @property + def supported_channels(self) -> list[AudioChannels]: + """Return a list of supported channels.""" + return [AudioChannels.CHANNEL_MONO] + + async def async_process_audio_stream( + self, metadata: SpeechMetadata, stream: AsyncIterable[bytes] + ) -> SpeechResult: + """Process an audio stream to STT service. + + Args: + metadata: Metadata about the audio stream + stream: Async iterable of audio chunks (PCM 16-bit, 16kHz, mono) + + Returns: + SpeechResult with the transcribed text + """ + _LOGGER.debug( + "开始处理音频流: language=%s, format=%s, codec=%s, sample_rate=%s", + metadata.language, + metadata.format, + metadata.codec, + metadata.sample_rate, + ) + + # 创建 ASR 配置 + config = ASRConfig( + credential_path=self._credential_path, + enable_punctuation=self._enable_punctuation, + sample_rate=16000, # HA 固定使用 16kHz + channels=1, # HA 固定使用单声道 + ) + + try: + # 使用 DoubaoASR 进行实时识别 + final_text = "" + async with DoubaoASR(config) as asr: + async for response in asr.transcribe_realtime(stream): + if response.type == ResponseType.FINAL_RESULT: + final_text = response.text + _LOGGER.debug("收到最终识别结果: %s", final_text) + elif response.type == ResponseType.INTERIM_RESULT: + _LOGGER.debug("收到中间识别结果: %s", response.text) + elif response.type == ResponseType.ERROR: + _LOGGER.error("识别过程出错: %s", response.error_msg) + return SpeechResult( + text=None, + result=SpeechResultState.ERROR, + ) + + if not final_text: + _LOGGER.warning("识别完成但未获得最终结果") + return SpeechResult( + text=None, + result=SpeechResultState.ERROR, + ) + + _LOGGER.info("语音识别成功: %s", final_text) + return SpeechResult( + text=final_text, + result=SpeechResultState.SUCCESS, + ) + + except ASRError as err: + _LOGGER.error("Doubao ASR 识别失败: %s", err) + return SpeechResult( + text=None, + result=SpeechResultState.ERROR, + ) + except Exception as err: # pylint: disable=broad-except + _LOGGER.exception("处理音频流时发生未知错误") + return SpeechResult( + text=None, + result=SpeechResultState.ERROR, + ) diff --git a/custom_components/hass_stt_doubao/translations/zh.json b/custom_components/hass_stt_doubao/translations/zh.json new file mode 100644 index 0000000..b0b6917 --- /dev/null +++ b/custom_components/hass_stt_doubao/translations/zh.json @@ -0,0 +1,33 @@ +{ + "config": { + "step": { + "user": { + "title": "配置 Doubao 语音识别", + "description": "设置 Doubao STT 集成", + "data": { + "credential_path": "凭据文件路径", + "enable_punctuation": "启用标点符号" + } + } + }, + "error": { + "cannot_connect": "无法连接到 Doubao 服务,请检查网络连接或凭据配置", + "unknown": "发生未知错误" + }, + "abort": { + "already_configured": "该设备已经配置过了" + } + }, + "options": { + "step": { + "init": { + "title": "Doubao STT 选项", + "description": "修改 Doubao STT 配置选项", + "data": { + "credential_path": "凭据文件路径", + "enable_punctuation": "启用标点符号" + } + } + } + } +} diff --git a/hacs.json b/hacs.json new file mode 100644 index 0000000..77d9946 --- /dev/null +++ b/hacs.json @@ -0,0 +1,6 @@ +{ + "name": "Doubao Speech to Text", + "render_readme": true, + "domains": ["stt"], + "country": ["CN"] +}