```
feat(stt): 新增豆包语音识别集成 - 实现完整的 Doubao STT 组件,支持中文语音识别 - 提供 UI 配置界面,无需手动编辑配置文件 - 支持实时流式识别和自动标点符号添加 - 集成自动设备注册和凭据管理功能 - 添加 HACS 和 hassfest 验证工作流 - 配置 .gitignore 和 Apache 2.0 许可证文件 - 完善 README 文档和故障排除指南 ```
This commit is contained in:
@@ -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"
|
||||
@@ -0,0 +1,3 @@
|
||||
.idea/
|
||||
*.mp3
|
||||
*.pyc
|
||||
@@ -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.
|
||||
@@ -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 配置支持
|
||||
- 自动设备注册
|
||||
- 中文语音识别
|
||||
@@ -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
|
||||
@@ -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."""
|
||||
@@ -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"
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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: ...
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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)"
|
||||
@@ -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
|
||||
@@ -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"
|
||||
}
|
||||
@@ -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
|
||||
@@ -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": "启用标点符号"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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": "启用标点符号"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user