mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-27 09:33:55 +08:00
合并main分支
This commit is contained in:
@@ -1,18 +1,22 @@
|
||||
from config.logger import setup_logging
|
||||
import time
|
||||
import copy
|
||||
from core.utils.util import remove_punctuation_and_length
|
||||
from core.handle.sendAudioHandle import send_stt_message
|
||||
from core.handle.intentHandler import handle_user_intent
|
||||
from core.utils.output_counter import check_device_output_limit
|
||||
from core.handle.reportHandle import enqueue_asr_report
|
||||
from core.utils.util import audio_to_data
|
||||
|
||||
TAG = __name__
|
||||
logger = setup_logging()
|
||||
|
||||
|
||||
async def handleAudioMessage(conn, audio):
|
||||
if not conn.asr_server_receive:
|
||||
logger.bind(tag=TAG).debug(f"前期数据处理中,暂停接收")
|
||||
if conn.vad is None:
|
||||
return
|
||||
if conn.client_listen_mode == "auto":
|
||||
if not conn.asr_server_receive:
|
||||
conn.logger.bind(tag=TAG).debug(f"前期数据处理中,暂停接收")
|
||||
return
|
||||
if conn.client_listen_mode == "auto" or conn.client_listen_mode == "realtime":
|
||||
have_voice = conn.vad.is_vad(conn, audio)
|
||||
else:
|
||||
have_voice = conn.client_have_voice
|
||||
@@ -35,12 +39,13 @@ async def handleAudioMessage(conn, audio):
|
||||
if len(conn.asr_audio) < 15:
|
||||
conn.asr_server_receive = True
|
||||
else:
|
||||
text, file_path = await conn.asr.speech_to_text(
|
||||
conn.asr_audio, conn.session_id
|
||||
)
|
||||
logger.bind(tag=TAG).info(f"识别文本: {text}")
|
||||
text, _ = await conn.asr.speech_to_text(conn.asr_audio, conn.session_id)
|
||||
conn.logger.bind(tag=TAG).info(f"识别文本: {text}")
|
||||
text_len, _ = remove_punctuation_and_length(text)
|
||||
if text_len > 0:
|
||||
# 使用自定义模块进行上报
|
||||
enqueue_asr_report(conn, text, copy.deepcopy(conn.asr_audio))
|
||||
|
||||
await startToChat(conn, text)
|
||||
else:
|
||||
conn.asr_server_receive = True
|
||||
@@ -49,6 +54,18 @@ async def handleAudioMessage(conn, audio):
|
||||
|
||||
|
||||
async def startToChat(conn, text):
|
||||
if conn.need_bind:
|
||||
await check_bind_device(conn)
|
||||
return
|
||||
|
||||
# 如果当日的输出字数大于限定的字数
|
||||
if conn.max_output_size > 0:
|
||||
if check_device_output_limit(
|
||||
conn.headers.get("device-id"), conn.max_output_size
|
||||
):
|
||||
await max_out_size(conn)
|
||||
return
|
||||
|
||||
# 首先进行意图分析
|
||||
intent_handled = await handle_user_intent(conn, text)
|
||||
|
||||
@@ -59,7 +76,7 @@ async def startToChat(conn, text):
|
||||
|
||||
# 意图未被处理,继续常规聊天流程
|
||||
await send_stt_message(conn, text)
|
||||
if conn.use_function_call_mode:
|
||||
if conn.intent_type == "function_call":
|
||||
# 使用支持function calling的聊天方法
|
||||
conn.executor.submit(conn.chat_with_function_calling, text)
|
||||
else:
|
||||
@@ -71,8 +88,8 @@ async def no_voice_close_connect(conn):
|
||||
conn.client_no_voice_last_time = time.time() * 1000
|
||||
else:
|
||||
no_voice_time = time.time() * 1000 - conn.client_no_voice_last_time
|
||||
close_connection_no_voice_time = conn.config.get(
|
||||
"close_connection_no_voice_time", 120
|
||||
close_connection_no_voice_time = int(
|
||||
conn.config.get("close_connection_no_voice_time", 120)
|
||||
)
|
||||
if (
|
||||
not conn.close_after_chat
|
||||
@@ -81,7 +98,65 @@ async def no_voice_close_connect(conn):
|
||||
conn.close_after_chat = True
|
||||
conn.client_abort = False
|
||||
conn.asr_server_receive = False
|
||||
prompt = (
|
||||
"请你以“时间过得真快”未来头,用富有感情、依依不舍的话来结束这场对话吧。"
|
||||
)
|
||||
end_prompt = conn.config.get("end_prompt", {})
|
||||
if end_prompt and end_prompt.get("enable", True) is False:
|
||||
conn.logger.bind(tag=TAG).info("结束对话,无需发送结束提示语")
|
||||
await conn.close()
|
||||
return
|
||||
prompt = end_prompt.get("prompt")
|
||||
if not prompt:
|
||||
prompt = "请你以“时间过得真快”未来头,用富有感情、依依不舍的话来结束这场对话吧。!"
|
||||
await startToChat(conn, prompt)
|
||||
|
||||
|
||||
async def max_out_size(conn):
|
||||
text = "不好意思,我现在有点事情要忙,明天这个时候我们再聊,约好了哦!明天不见不散,拜拜!"
|
||||
await send_stt_message(conn, text)
|
||||
conn.tts_first_text_index = 0
|
||||
conn.tts_last_text_index = 0
|
||||
conn.llm_finish_task = True
|
||||
file_path = "config/assets/max_output_size.wav"
|
||||
opus_packets, _ = audio_to_data(file_path)
|
||||
conn.audio_play_queue.put((opus_packets, text, 0))
|
||||
conn.close_after_chat = True
|
||||
|
||||
|
||||
async def check_bind_device(conn):
|
||||
if conn.bind_code:
|
||||
# 确保bind_code是6位数字
|
||||
if len(conn.bind_code) != 6:
|
||||
conn.logger.bind(tag=TAG).error(f"无效的绑定码格式: {conn.bind_code}")
|
||||
text = "绑定码格式错误,请检查配置。"
|
||||
await send_stt_message(conn, text)
|
||||
return
|
||||
|
||||
text = f"请登录控制面板,输入{conn.bind_code},绑定设备。"
|
||||
await send_stt_message(conn, text)
|
||||
conn.tts_first_text_index = 0
|
||||
conn.tts_last_text_index = 6
|
||||
conn.llm_finish_task = True
|
||||
|
||||
# 播放提示音
|
||||
music_path = "config/assets/bind_code.wav"
|
||||
opus_packets, _ = audio_to_data(music_path)
|
||||
conn.audio_play_queue.put((opus_packets, text, 0))
|
||||
|
||||
# 逐个播放数字
|
||||
for i in range(6): # 确保只播放6位数字
|
||||
try:
|
||||
digit = conn.bind_code[i]
|
||||
num_path = f"config/assets/bind_code/{digit}.wav"
|
||||
num_packets, _ = audio_to_data(num_path)
|
||||
conn.audio_play_queue.put((num_packets, None, i + 1))
|
||||
except Exception as e:
|
||||
conn.logger.bind(tag=TAG).error(f"播放数字音频失败: {e}")
|
||||
continue
|
||||
else:
|
||||
text = f"没有找到该设备的版本信息,请正确配置 OTA地址,然后重新编译固件。"
|
||||
await send_stt_message(conn, text)
|
||||
conn.tts_first_text_index = 0
|
||||
conn.tts_last_text_index = 0
|
||||
conn.llm_finish_task = True
|
||||
music_path = "config/assets/bind_not_found.wav"
|
||||
opus_packets, _ = audio_to_data(music_path)
|
||||
conn.audio_play_queue.put((opus_packets, text, 0))
|
||||
|
||||
@@ -10,40 +10,62 @@ from pydub import AudioSegment
|
||||
from core.providers.tts.dto.dto import TTSMessageDTO, MsgType, SentenceType
|
||||
from core.utils.util import check_model_key
|
||||
from core.providers.tts.base import TTSProviderBase
|
||||
from config.logger import setup_logging
|
||||
|
||||
TAG = __name__
|
||||
logger = setup_logging()
|
||||
|
||||
|
||||
class TTSProvider(TTSProviderBase):
|
||||
def __init__(self, config, delete_audio_file):
|
||||
super().__init__(config, delete_audio_file)
|
||||
self.appid = config.get("appid")
|
||||
if config.get("appid"):
|
||||
self.appid = int(config.get("appid"))
|
||||
else:
|
||||
self.appid = ""
|
||||
self.access_token = config.get("access_token")
|
||||
self.cluster = config.get("cluster")
|
||||
self.voice = config.get("voice")
|
||||
|
||||
if config.get("private_voice"):
|
||||
self.voice = config.get("private_voice")
|
||||
else:
|
||||
self.voice = config.get("voice")
|
||||
|
||||
# 处理空字符串的情况
|
||||
speed_ratio = config.get("speed_ratio", "1.0")
|
||||
volume_ratio = config.get("volume_ratio", "1.0")
|
||||
pitch_ratio = config.get("pitch_ratio", "1.0")
|
||||
|
||||
self.speed_ratio = float(speed_ratio) if speed_ratio else 1.0
|
||||
self.volume_ratio = float(volume_ratio) if volume_ratio else 1.0
|
||||
self.pitch_ratio = float(pitch_ratio) if pitch_ratio else 1.0
|
||||
|
||||
self.api_url = config.get("api_url")
|
||||
self.authorization = config.get("authorization")
|
||||
self.header = {"Authorization": f"{self.authorization}{self.access_token}"}
|
||||
check_model_key("TTS", self.access_token)
|
||||
|
||||
def generate_filename(self, extension=".wav"):
|
||||
return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}")
|
||||
return os.path.join(
|
||||
self.output_file,
|
||||
f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}",
|
||||
)
|
||||
|
||||
async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False):
|
||||
tmp_file = self.generate_filename()
|
||||
request_json = {
|
||||
"app": {
|
||||
"appid": f"{self.appid}",
|
||||
"token": "access_token",
|
||||
"cluster": self.cluster
|
||||
},
|
||||
"user": {
|
||||
"uid": "1"
|
||||
"token": self.access_token,
|
||||
"cluster": self.cluster,
|
||||
},
|
||||
"user": {"uid": "1"},
|
||||
"audio": {
|
||||
"voice_type": self.voice,
|
||||
"encoding": "wav",
|
||||
"speed_ratio": 1.0,
|
||||
"volume_ratio": 1.0,
|
||||
"pitch_ratio": 1.0,
|
||||
"speed_ratio": self.speed_ratio,
|
||||
"volume_ratio": self.volume_ratio,
|
||||
"pitch_ratio": self.pitch_ratio,
|
||||
},
|
||||
"request": {
|
||||
"reqid": str(uuid.uuid4()),
|
||||
@@ -51,18 +73,22 @@ class TTSProvider(TTSProviderBase):
|
||||
"text_type": "plain",
|
||||
"operation": "query",
|
||||
"with_frontend": 1,
|
||||
"frontend_type": "unitTson"
|
||||
}
|
||||
"frontend_type": "unitTson",
|
||||
},
|
||||
}
|
||||
|
||||
try:
|
||||
resp = requests.post(self.api_url, json.dumps(request_json), headers=self.header)
|
||||
resp = requests.post(
|
||||
self.api_url, json.dumps(request_json), headers=self.header
|
||||
)
|
||||
if "data" in resp.json():
|
||||
data = resp.json()["data"]
|
||||
file_to_save = open(tmp_file, "wb")
|
||||
file_to_save.write(base64.b64decode(data))
|
||||
else:
|
||||
raise Exception(f"{__name__} status_code: {resp.status_code} response: {resp.content}")
|
||||
raise Exception(
|
||||
f"{__name__} status_code: {resp.status_code} response: {resp.content}"
|
||||
)
|
||||
except Exception as e:
|
||||
raise Exception(f"{__name__} error: {e}")
|
||||
# 使用 pydub 读取临时文件
|
||||
|
||||
@@ -16,7 +16,10 @@ class TTSProvider(TTSProviderBase):
|
||||
self.appid = config.get("appid")
|
||||
self.secret_id = config.get("secret_id")
|
||||
self.secret_key = config.get("secret_key")
|
||||
self.voice = config.get("voice")
|
||||
if config.get("private_voice"):
|
||||
self.voice = config.get("private_voice")
|
||||
else:
|
||||
self.voice = int(config.get("voice"))
|
||||
self.api_url = "https://tts.tencentcloudapi.com" # 正确的API端点
|
||||
self.region = config.get("region")
|
||||
self.output_file = config.get("output_dir")
|
||||
@@ -25,35 +28,36 @@ class TTSProvider(TTSProviderBase):
|
||||
"""生成鉴权请求头"""
|
||||
# 获取当前UTC时间戳
|
||||
timestamp = int(time.time())
|
||||
|
||||
|
||||
# 使用UTC时间计算日期
|
||||
utc_date = datetime.fromtimestamp(timestamp, tz=timezone.utc).strftime('%Y-%m-%d')
|
||||
|
||||
utc_date = datetime.fromtimestamp(timestamp, tz=timezone.utc).strftime(
|
||||
"%Y-%m-%d"
|
||||
)
|
||||
|
||||
# 服务名称必须是 "tts"
|
||||
service = "tts"
|
||||
|
||||
|
||||
# 拼接凭证范围
|
||||
credential_scope = f"{utc_date}/{service}/tc3_request"
|
||||
|
||||
|
||||
# 使用TC3-HMAC-SHA256签名方法
|
||||
algorithm = "TC3-HMAC-SHA256"
|
||||
|
||||
|
||||
# 构建规范请求字符串
|
||||
http_request_method = "POST"
|
||||
canonical_uri = "/"
|
||||
canonical_querystring = ""
|
||||
|
||||
|
||||
# 请求头必须包含host和content-type,且按字典序排列
|
||||
canonical_headers = (
|
||||
f"content-type:application/json\n"
|
||||
f"host:tts.tencentcloudapi.com\n"
|
||||
f"content-type:application/json\n" f"host:tts.tencentcloudapi.com\n"
|
||||
)
|
||||
signed_headers = "content-type;host"
|
||||
|
||||
|
||||
# 请求体哈希值
|
||||
payload = json.dumps(request_body)
|
||||
payload_hash = hashlib.sha256(payload.encode('utf-8')).hexdigest()
|
||||
|
||||
payload_hash = hashlib.sha256(payload.encode("utf-8")).hexdigest()
|
||||
|
||||
# 构建规范请求字符串
|
||||
canonical_request = (
|
||||
f"{http_request_method}\n"
|
||||
@@ -63,10 +67,12 @@ class TTSProvider(TTSProviderBase):
|
||||
f"{signed_headers}\n"
|
||||
f"{payload_hash}"
|
||||
)
|
||||
|
||||
|
||||
# 计算规范请求的哈希值
|
||||
hashed_canonical_request = hashlib.sha256(canonical_request.encode('utf-8')).hexdigest()
|
||||
|
||||
hashed_canonical_request = hashlib.sha256(
|
||||
canonical_request.encode("utf-8")
|
||||
).hexdigest()
|
||||
|
||||
# 构建待签名字符串
|
||||
string_to_sign = (
|
||||
f"{algorithm}\n"
|
||||
@@ -74,19 +80,19 @@ class TTSProvider(TTSProviderBase):
|
||||
f"{credential_scope}\n"
|
||||
f"{hashed_canonical_request}"
|
||||
)
|
||||
|
||||
|
||||
# 计算签名密钥
|
||||
secret_date = self._hmac_sha256(f"TC3{self.secret_key}".encode('utf-8'), utc_date)
|
||||
secret_date = self._hmac_sha256(
|
||||
f"TC3{self.secret_key}".encode("utf-8"), utc_date
|
||||
)
|
||||
secret_service = self._hmac_sha256(secret_date, service)
|
||||
secret_signing = self._hmac_sha256(secret_service, "tc3_request")
|
||||
|
||||
|
||||
# 计算签名
|
||||
signature = hmac.new(
|
||||
secret_signing,
|
||||
string_to_sign.encode('utf-8'),
|
||||
hashlib.sha256
|
||||
secret_signing, string_to_sign.encode("utf-8"), hashlib.sha256
|
||||
).hexdigest()
|
||||
|
||||
|
||||
# 构建授权头
|
||||
authorization = (
|
||||
f"{algorithm} "
|
||||
@@ -94,7 +100,7 @@ class TTSProvider(TTSProviderBase):
|
||||
f"SignedHeaders={signed_headers}, "
|
||||
f"Signature={signature}"
|
||||
)
|
||||
|
||||
|
||||
# 构建请求头
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
@@ -104,19 +110,22 @@ class TTSProvider(TTSProviderBase):
|
||||
"X-TC-Timestamp": str(timestamp),
|
||||
"X-TC-Version": "2019-08-23",
|
||||
"X-TC-Region": self.region,
|
||||
"X-TC-Language": "zh-CN"
|
||||
"X-TC-Language": "zh-CN",
|
||||
}
|
||||
|
||||
|
||||
return headers
|
||||
|
||||
|
||||
def _hmac_sha256(self, key, msg):
|
||||
"""HMAC-SHA256加密"""
|
||||
if isinstance(msg, str):
|
||||
msg = msg.encode('utf-8')
|
||||
msg = msg.encode("utf-8")
|
||||
return hmac.new(key, msg, hashlib.sha256).digest()
|
||||
|
||||
def generate_filename(self, extension=".wav"):
|
||||
return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}")
|
||||
return os.path.join(
|
||||
self.output_file,
|
||||
f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}",
|
||||
)
|
||||
|
||||
async def text_to_speak(self, text, output_file):
|
||||
# 构建请求体
|
||||
@@ -129,19 +138,23 @@ class TTSProvider(TTSProviderBase):
|
||||
try:
|
||||
# 获取请求头(每次请求都重新生成,以确保时间戳和签名是最新的)
|
||||
headers = self._get_auth_headers(request_json)
|
||||
|
||||
|
||||
# 发送请求
|
||||
resp = requests.post(self.api_url, json.dumps(request_json), headers=headers)
|
||||
|
||||
resp = requests.post(
|
||||
self.api_url, json.dumps(request_json), headers=headers
|
||||
)
|
||||
|
||||
# 检查响应
|
||||
if resp.status_code == 200:
|
||||
response_data = resp.json()
|
||||
|
||||
|
||||
# 检查是否成功
|
||||
if response_data.get("Response", {}).get("Error") is not None:
|
||||
error_info = response_data["Response"]["Error"]
|
||||
raise Exception(f"API返回错误: {error_info['Code']}: {error_info['Message']}")
|
||||
|
||||
raise Exception(
|
||||
f"API返回错误: {error_info['Code']}: {error_info['Message']}"
|
||||
)
|
||||
|
||||
# 提取音频数据
|
||||
audio_data = response_data["Response"].get("Audio")
|
||||
if audio_data:
|
||||
@@ -151,6 +164,8 @@ class TTSProvider(TTSProviderBase):
|
||||
else:
|
||||
raise Exception(f"{__name__}: 没有返回音频数据: {response_data}")
|
||||
else:
|
||||
raise Exception(f"{__name__} status_code: {resp.status_code} response: {resp.content}")
|
||||
raise Exception(
|
||||
f"{__name__} status_code: {resp.status_code} response: {resp.content}"
|
||||
)
|
||||
except Exception as e:
|
||||
raise Exception(f"{__name__} error: {e}")
|
||||
raise Exception(f"{__name__} error: {e}")
|
||||
|
||||
Reference in New Issue
Block a user