update: 更改to_tts保存临时文件判断

This commit is contained in:
Sakura-RanChen
2025-06-04 16:46:49 +08:00
parent 3657f6ce75
commit 23cb7616d9
17 changed files with 253 additions and 92 deletions
+15 -17
View File
@@ -2,10 +2,9 @@ import os
import time import time
import json import json
import random import random
import shutil
import asyncio import asyncio
from core.handle.sendAudioHandle import send_stt_message from core.handle.sendAudioHandle import send_stt_message
from core.utils.util import remove_punctuation_and_length from core.utils.util import remove_punctuation_and_length, opus_datas_to_wav_bytes
from core.providers.tts.dto.dto import ContentType, InterfaceType from core.providers.tts.dto.dto import ContentType, InterfaceType
from core.handle.mcpHandle import ( from core.handle.mcpHandle import (
MCPClient, MCPClient,
@@ -119,19 +118,18 @@ async def wakeupWordsResponse(conn):
result = conn.llm.response_no_stream(conn.config["prompt"], question) result = conn.llm.response_no_stream(conn.config["prompt"], question)
if result is None or result == "": if result is None or result == "":
return return
tts_file = await asyncio.to_thread(conn.tts.to_tts, result)
if tts_file is not None and os.path.exists(tts_file): opus_datas = await asyncio.to_thread(conn.tts.to_tts, result)
file_type = os.path.splitext(tts_file)[1] if not opus_datas:
if file_type: return
file_type = file_type.lstrip(".")
old_file = getWakeupWordFile("my_" + WAKEUP_CONFIG["file_name"]) wav_bytes = opus_datas_to_wav_bytes(opus_datas, sample_rate=16000)
if old_file is not None: file_path = os.path.join(
os.remove(old_file) WAKEUP_CONFIG["dir"], "my_" + WAKEUP_CONFIG["file_name"] + ".wav"
"""将文件挪到"wakeup_words.mp3""" )
shutil.move( # 写入wav数据
tts_file, with open(file_path, "wb") as f:
WAKEUP_CONFIG["dir"] + "my_" + WAKEUP_CONFIG["file_name"] + "." + file_type, f.write(wav_bytes)
)
WAKEUP_CONFIG["create_time"] = time.time() WAKEUP_CONFIG["create_time"] = time.time()
WAKEUP_CONFIG["text"] = result WAKEUP_CONFIG["text"] = result
@@ -91,7 +91,7 @@ class TTSProvider(TTSProviderBase):
self.appkey = config.get("appkey") self.appkey = config.get("appkey")
self.format = config.get("format", "wav") self.format = config.get("format", "wav")
self.audio_file_type = config.get("format", "wav")
sample_rate = config.get("sample_rate", "16000") sample_rate = config.get("sample_rate", "16000")
self.sample_rate = int(sample_rate) if sample_rate else 16000 self.sample_rate = int(sample_rate) if sample_rate else 16000
@@ -188,9 +188,12 @@ class TTSProvider(TTSProviderBase):
) )
# 检查返回请求数据的mime类型是否是audio/***,是则保存到指定路径下;返回的是binary格式的 # 检查返回请求数据的mime类型是否是audio/***,是则保存到指定路径下;返回的是binary格式的
if resp.headers["Content-Type"].startswith("audio/"): if resp.headers["Content-Type"].startswith("audio/"):
with open(output_file, "wb") as f: if output_file:
f.write(resp.content) with open(output_file, "wb") as f:
return output_file f.write(resp.content)
return output_file
else:
return resp.content
else: else:
raise Exception( raise Exception(
f"{__name__} status_code: {resp.status_code} response: {resp.content}" f"{__name__} status_code: {resp.status_code} response: {resp.content}"
+67 -28
View File
@@ -8,7 +8,7 @@ from datetime import datetime
from core.utils import textUtils from core.utils import textUtils
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from config.logger import setup_logging from config.logger import setup_logging
from core.utils.util import audio_to_data from core.utils.util import audio_to_data, audio_bytes_to_data
from core.utils.tts import MarkdownCleaner from core.utils.tts import MarkdownCleaner
from core.utils.output_counter import add_device_output from core.utils.output_counter import add_device_output
from core.handle.reportHandle import enqueue_tts_report from core.handle.reportHandle import enqueue_tts_report
@@ -20,7 +20,6 @@ from core.providers.tts.dto.dto import (
InterfaceType, InterfaceType,
) )
import traceback import traceback
TAG = __name__ TAG = __name__
@@ -33,6 +32,7 @@ class TTSProviderBase(ABC):
self.conn = None self.conn = None
self.tts_timeout = 10 self.tts_timeout = 10
self.delete_audio_file = delete_audio_file self.delete_audio_file = delete_audio_file
self.audio_file_type = "wav"
self.output_file = config.get("output_dir", "tmp/") self.output_file = config.get("output_dir", "tmp/")
self.tts_text_queue = queue.Queue() self.tts_text_queue = queue.Queue()
self.tts_audio_queue = queue.Queue() self.tts_audio_queue = queue.Queue()
@@ -76,35 +76,60 @@ class TTSProviderBase(ABC):
) )
def to_tts(self, text): def to_tts(self, text):
tmp_file = self.generate_filename() text = MarkdownCleaner.clean_markdown(text)
try: max_repeat_time = 5
max_repeat_time = 5 if self.delete_audio_file:
text = MarkdownCleaner.clean_markdown(text) # 需要删除文件的直接转为音频数据
while not os.path.exists(tmp_file) and max_repeat_time > 0: while max_repeat_time > 0:
try: try:
asyncio.run(self.text_to_speak(text, tmp_file)) audio_bytes = asyncio.run(self.text_to_speak(text, None))
if audio_bytes:
audio_datas, _ = audio_bytes_to_data(audio_bytes, file_type=self.audio_file_type, is_opus=True)
return audio_datas
else:
max_repeat_time -= 1
except Exception as e: except Exception as e:
logger.bind(tag=TAG).warning( logger.bind(tag=TAG).warning(
f"语音生成失败{5 - max_repeat_time + 1}次: {text},错误: {e}" f"语音生成失败{5 - max_repeat_time + 1}次: {text},错误: {e}"
) )
# 未执行成功,删除文件
if os.path.exists(tmp_file):
os.remove(tmp_file)
max_repeat_time -= 1 max_repeat_time -= 1
if max_repeat_time > 0: if max_repeat_time > 0:
logger.bind(tag=TAG).info( logger.bind(tag=TAG).info(
f"语音生成成功: {text}:{tmp_file},重试{5 - max_repeat_time}" f"语音生成成功: {text},重试{5 - max_repeat_time}"
) )
else: else:
logger.bind(tag=TAG).error( logger.bind(tag=TAG).error(
f"语音生成失败: {text},请检查网络或服务是否正常" f"语音生成失败: {text},请检查网络或服务是否正常"
) )
return tmp_file
except Exception as e:
logger.bind(tag=TAG).error(f"Failed to generate TTS file: {e}")
return None return None
else:
tmp_file = self.generate_filename()
try:
while not os.path.exists(tmp_file) and max_repeat_time > 0:
try:
asyncio.run(self.text_to_speak(text, tmp_file))
except Exception as e:
logger.bind(tag=TAG).warning(
f"语音生成失败{5 - max_repeat_time + 1}次: {text},错误: {e}"
)
# 未执行成功,删除文件
if os.path.exists(tmp_file):
os.remove(tmp_file)
max_repeat_time -= 1
if max_repeat_time > 0:
logger.bind(tag=TAG).info(
f"语音生成成功: {text}:{tmp_file},重试{5 - max_repeat_time}"
)
else:
logger.bind(tag=TAG).error(
f"语音生成失败: {text},请检查网络或服务是否正常"
)
return tmp_file
except Exception as e:
logger.bind(tag=TAG).error(f"Failed to generate TTS file: {e}")
return None
@abstractmethod @abstractmethod
async def text_to_speak(self, text, output_file): async def text_to_speak(self, text, output_file):
@@ -192,12 +217,19 @@ class TTSProviderBase(ABC):
self.tts_text_buff.append(message.content_detail) self.tts_text_buff.append(message.content_detail)
segment_text = self._get_segment_text() segment_text = self._get_segment_text()
if segment_text: if segment_text:
tts_file = self.to_tts(segment_text) if self.delete_audio_file:
if tts_file: audio_datas = self.to_tts(segment_text)
audio_datas = self._process_audio_file(tts_file) if audio_datas:
self.tts_audio_queue.put( self.tts_audio_queue.put(
(message.sentence_type, audio_datas, segment_text) (message.sentence_type, audio_datas, segment_text)
) )
else:
tts_file = self.to_tts(segment_text)
if tts_file:
audio_datas = self._process_audio_file(tts_file)
self.tts_audio_queue.put(
(message.sentence_type, audio_datas, segment_text)
)
elif ContentType.FILE == message.content_type: elif ContentType.FILE == message.content_type:
self._process_remaining_text() self._process_remaining_text()
tts_file = message.content_file tts_file = message.content_file
@@ -334,11 +366,18 @@ class TTSProviderBase(ABC):
if remaining_text: if remaining_text:
segment_text = textUtils.get_string_no_punctuation_or_emoji(remaining_text) segment_text = textUtils.get_string_no_punctuation_or_emoji(remaining_text)
if segment_text: if segment_text:
tts_file = self.to_tts(segment_text) if self.delete_audio_file:
audio_datas = self._process_audio_file(tts_file) audio_datas = self.to_tts(segment_text)
self.tts_audio_queue.put( if audio_datas:
(SentenceType.MIDDLE, audio_datas, segment_text) self.tts_audio_queue.put(
) (SentenceType.MIDDLE, audio_datas, segment_text)
)
else:
tts_file = self.to_tts(segment_text)
audio_datas = self._process_audio_file(tts_file)
self.tts_audio_queue.put(
(SentenceType.MIDDLE, audio_datas, segment_text)
)
self.processed_chars += len(full_text) self.processed_chars += len(full_text)
return True return True
return False return False
@@ -11,8 +11,8 @@ class TTSProvider(TTSProviderBase):
self.voice = config.get("private_voice") self.voice = config.get("private_voice")
else: else:
self.voice = config.get("voice") self.voice = config.get("voice")
self.response_format = config.get("response_format") self.response_format = config.get("response_format", "mp3")
self.audio_file_type = config.get("response_format", "mp3")
self.host = "api.coze.cn" self.host = "api.coze.cn"
self.api_url = f"https://{self.host}/v1/audio/speech" self.api_url = f"https://{self.host}/v1/audio/speech"
@@ -33,7 +33,10 @@ class TTSProvider(TTSProviderBase):
"POST", self.api_url, json=request_json, headers=headers "POST", self.api_url, json=request_json, headers=headers
) )
data = response.content data = response.content
file_to_save = open(output_file, "wb") if output_file:
file_to_save.write(data) with open(output_file, "wb") as file_to_save:
file_to_save.write(data)
else:
return data
except Exception as e: except Exception as e:
raise Exception(f"{__name__} error: {e}") raise Exception(f"{__name__} error: {e}")
@@ -16,8 +16,8 @@ class TTSProvider(TTSProviderBase):
self.method = config.get("method", "GET") self.method = config.get("method", "GET")
self.headers = config.get("headers", {}) self.headers = config.get("headers", {})
self.format = config.get("format", "wav") self.format = config.get("format", "wav")
self.audio_file_type = config.get("format", "wav")
self.output_file = config.get("output_dir", "tmp/") self.output_file = config.get("output_dir", "tmp/")
self.params = config.get("params") self.params = config.get("params")
if isinstance(self.params, str): if isinstance(self.params, str):
@@ -43,8 +43,11 @@ class TTSProvider(TTSProviderBase):
else: else:
resp = requests.get(self.url, params=request_params, headers=self.headers) resp = requests.get(self.url, params=request_params, headers=self.headers)
if resp.status_code == 200: if resp.status_code == 200:
with open(output_file, "wb") as file: if output_file:
file.write(resp.content) with open(output_file, "wb") as file:
file.write(resp.content)
else:
return resp.content
else: else:
error_msg = f"Custom TTS请求失败: {resp.status_code} - {resp.text}" error_msg = f"Custom TTS请求失败: {resp.status_code} - {resp.text}"
logger.bind(tag=TAG).error(error_msg) logger.bind(tag=TAG).error(error_msg)
@@ -29,7 +29,7 @@ class TTSProvider(TTSProviderBase):
speed_ratio = config.get("speed_ratio", "1.0") speed_ratio = config.get("speed_ratio", "1.0")
volume_ratio = config.get("volume_ratio", "1.0") volume_ratio = config.get("volume_ratio", "1.0")
pitch_ratio = config.get("pitch_ratio", "1.0") pitch_ratio = config.get("pitch_ratio", "1.0")
self.audio_file_type = config.get("format", "wav")
self.speed_ratio = float(speed_ratio) if speed_ratio else 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.volume_ratio = float(volume_ratio) if volume_ratio else 1.0
self.pitch_ratio = float(pitch_ratio) if pitch_ratio else 1.0 self.pitch_ratio = float(pitch_ratio) if pitch_ratio else 1.0
@@ -49,7 +49,7 @@ class TTSProvider(TTSProviderBase):
"user": {"uid": "1"}, "user": {"uid": "1"},
"audio": { "audio": {
"voice_type": self.voice, "voice_type": self.voice,
"encoding": "wav", "encoding": self.audio_file_type,
"speed_ratio": self.speed_ratio, "speed_ratio": self.speed_ratio,
"volume_ratio": self.volume_ratio, "volume_ratio": self.volume_ratio,
"pitch_ratio": self.pitch_ratio, "pitch_ratio": self.pitch_ratio,
@@ -70,8 +70,12 @@ class TTSProvider(TTSProviderBase):
) )
if "data" in resp.json(): if "data" in resp.json():
data = resp.json()["data"] data = resp.json()["data"]
file_to_save = open(output_file, "wb") audio_bytes = base64.b64decode(data)
file_to_save.write(base64.b64decode(data)) if output_file:
with open(output_file, "wb") as file_to_save:
file_to_save.write(audio_bytes)
else:
return audio_bytes
else: else:
raise Exception( raise Exception(
f"{__name__} status_code: {resp.status_code} response: {resp.content}" f"{__name__} status_code: {resp.status_code} response: {resp.content}"
+17 -8
View File
@@ -12,6 +12,7 @@ class TTSProvider(TTSProviderBase):
self.voice = config.get("private_voice") self.voice = config.get("private_voice")
else: else:
self.voice = config.get("voice") self.voice = config.get("voice")
self.audio_file_type = config.get("format", "mp3")
def generate_filename(self, extension=".mp3"): def generate_filename(self, extension=".mp3"):
return os.path.join( return os.path.join(
@@ -22,16 +23,24 @@ class TTSProvider(TTSProviderBase):
async def text_to_speak(self, text, output_file): async def text_to_speak(self, text, output_file):
try: try:
communicate = edge_tts.Communicate(text, voice=self.voice) communicate = edge_tts.Communicate(text, voice=self.voice)
# 确保目录存在并创建空文件 if output_file:
os.makedirs(os.path.dirname(output_file), exist_ok=True) # 确保目录存在并创建空文件
with open(output_file, "wb") as f: os.makedirs(os.path.dirname(output_file), exist_ok=True)
pass with open(output_file, "wb") as f:
pass
# 流式写入音频数据 # 流式写入音频数据
with open(output_file, "ab") as f: # 改为追加模式避免覆盖 with open(output_file, "ab") as f: # 改为追加模式避免覆盖
async for chunk in communicate.stream():
if chunk["type"] == "audio": # 只处理音频数据块
f.write(chunk["data"])
else:
# 返回音频二进制数据
audio_bytes = b""
async for chunk in communicate.stream(): async for chunk in communicate.stream():
if chunk["type"] == "audio": # 只处理音频数据块 if chunk["type"] == "audio":
f.write(chunk["data"]) audio_bytes += chunk["data"]
return audio_bytes
except Exception as e: except Exception as e:
error_msg = f"Edge TTS请求失败: {e}" error_msg = f"Edge TTS请求失败: {e}"
raise Exception(error_msg) # 抛出异常,让调用方捕获 raise Exception(error_msg) # 抛出异常,让调用方捕获
@@ -88,7 +88,7 @@ class TTSProvider(TTSProviderBase):
self.reference_audio = parse_string_to_list(config.get("reference_audio")) self.reference_audio = parse_string_to_list(config.get("reference_audio"))
self.reference_text = parse_string_to_list(config.get("reference_text")) self.reference_text = parse_string_to_list(config.get("reference_text"))
self.format = config.get("response_format", "wav") self.format = config.get("response_format", "wav")
self.audio_file_type = config.get("response_format", "wav")
self.api_key = config.get("api_key", "YOUR_API_KEY") self.api_key = config.get("api_key", "YOUR_API_KEY")
have_key = check_model_key("FishSpeech TTS", self.api_key) have_key = check_model_key("FishSpeech TTS", self.api_key)
if not have_key: if not have_key:
@@ -170,8 +170,11 @@ class TTSProvider(TTSProviderBase):
if response.status_code == 200: if response.status_code == 200:
audio_content = response.content audio_content = response.content
with open(output_file, "wb") as audio_file: if output_file:
audio_file.write(audio_content) with open(output_file, "wb") as audio_file:
audio_file.write(audio_content)
else:
return audio_content
else: else:
error_msg = f"Request failed with status code {response.status_code}" error_msg = f"Request failed with status code {response.status_code}"
@@ -65,6 +65,7 @@ class TTSProvider(TTSProviderBase):
self.aux_ref_audio_paths = parse_string_to_list( self.aux_ref_audio_paths = parse_string_to_list(
config.get("aux_ref_audio_paths") config.get("aux_ref_audio_paths")
) )
self.audio_file_type = config.get("format", "wav")
async def text_to_speak(self, text, output_file): async def text_to_speak(self, text, output_file):
request_json = { request_json = {
@@ -91,8 +92,11 @@ class TTSProvider(TTSProviderBase):
resp = requests.post(self.url, json=request_json) resp = requests.post(self.url, json=request_json)
if resp.status_code == 200: if resp.status_code == 200:
with open(output_file, "wb") as file: if output_file:
file.write(resp.content) with open(output_file, "wb") as file:
file.write(resp.content)
else:
return resp.content
else: else:
error_msg = f"GPT_SoVITS_V2 TTS请求失败: {resp.status_code} - {resp.text}" error_msg = f"GPT_SoVITS_V2 TTS请求失败: {resp.status_code} - {resp.text}"
logger.bind(tag=TAG).error(error_msg) logger.bind(tag=TAG).error(error_msg)
@@ -32,6 +32,7 @@ class TTSProvider(TTSProviderBase):
self.cut_punc = config.get("cut_punc", "") self.cut_punc = config.get("cut_punc", "")
self.inp_refs = parse_string_to_list(config.get("inp_refs")) self.inp_refs = parse_string_to_list(config.get("inp_refs"))
self.if_sr = str(config.get("if_sr", False)).lower() in ("true", "1", "yes") self.if_sr = str(config.get("if_sr", False)).lower() in ("true", "1", "yes")
self.audio_file_type = config.get("format", "wav")
async def text_to_speak(self, text, output_file): async def text_to_speak(self, text, output_file):
request_params = { request_params = {
@@ -52,8 +53,11 @@ class TTSProvider(TTSProviderBase):
resp = requests.get(self.url, params=request_params) resp = requests.get(self.url, params=request_params)
if resp.status_code == 200: if resp.status_code == 200:
with open(output_file, "wb") as file: if output_file:
file.write(resp.content) with open(output_file, "wb") as file:
file.write(resp.content)
else:
return resp.content
else: else:
error_msg = f"GPT_SoVITS_V3 TTS请求失败: {resp.status_code} - {resp.text}" error_msg = f"GPT_SoVITS_V3 TTS请求失败: {resp.status_code} - {resp.text}"
logger.bind(tag=TAG).error(error_msg) logger.bind(tag=TAG).error(error_msg)
@@ -52,6 +52,7 @@ class TTSProvider(TTSProviderBase):
"Content-Type": "application/json", "Content-Type": "application/json",
"Authorization": f"Bearer {self.api_key}", "Authorization": f"Bearer {self.api_key}",
} }
self.audio_file_type = defult_audio_setting.get("format", "mp3")
def generate_filename(self, extension=".mp3"): def generate_filename(self, extension=".mp3"):
return os.path.join( return os.path.join(
@@ -80,8 +81,12 @@ class TTSProvider(TTSProviderBase):
# 检查返回请求数据的status_code是否为0 # 检查返回请求数据的status_code是否为0
if resp.json()["base_resp"]["status_code"] == 0: if resp.json()["base_resp"]["status_code"] == 0:
data = resp.json()["data"]["audio"] data = resp.json()["data"]["audio"]
file_to_save = open(output_file, "wb") audio_bytes = bytes.fromhex(data)
file_to_save.write(bytes.fromhex(data)) if output_file:
with open(output_file, "wb") as file_to_save:
file_to_save.write(audio_bytes)
else:
return audio_bytes
else: else:
raise Exception( raise Exception(
f"{__name__} status_code: {resp.status_code} response: {resp.content}" f"{__name__} status_code: {resp.status_code} response: {resp.content}"
@@ -17,7 +17,8 @@ class TTSProvider(TTSProviderBase):
self.voice = config.get("private_voice") self.voice = config.get("private_voice")
else: else:
self.voice = config.get("voice", "alloy") self.voice = config.get("voice", "alloy")
self.response_format = "wav" self.response_format = config.get("format", "wav")
self.audio_file_type = config.get("format", "wav")
# 处理空字符串的情况 # 处理空字符串的情况
speed = config.get("speed", "1.0") speed = config.get("speed", "1.0")
@@ -40,8 +41,11 @@ class TTSProvider(TTSProviderBase):
} }
response = requests.post(self.api_url, json=data, headers=headers) response = requests.post(self.api_url, json=data, headers=headers)
if response.status_code == 200: if response.status_code == 200:
with open(output_file, "wb") as audio_file: if output_file:
audio_file.write(response.content) with open(output_file, "wb") as audio_file:
audio_file.write(response.content)
else:
return response.content
else: else:
raise Exception( raise Exception(
f"OpenAI TTS请求失败: {response.status_code} - {response.text}" f"OpenAI TTS请求失败: {response.status_code} - {response.text}"
@@ -11,7 +11,8 @@ class TTSProvider(TTSProviderBase):
self.voice = config.get("private_voice") self.voice = config.get("private_voice")
else: else:
self.voice = config.get("voice") self.voice = config.get("voice")
self.response_format = config.get("response_format") self.response_format = config.get("response_format", "mp3")
self.audio_file_type = config.get("response_format", "mp3")
self.sample_rate = config.get("sample_rate") self.sample_rate = config.get("sample_rate")
self.speed = float(config.get("speed", 1.0)) self.speed = float(config.get("speed", 1.0))
self.gain = config.get("gain") self.gain = config.get("gain")
@@ -35,7 +36,10 @@ class TTSProvider(TTSProviderBase):
"POST", self.api_url, json=request_json, headers=headers "POST", self.api_url, json=request_json, headers=headers
) )
data = response.content data = response.content
file_to_save = open(output_file, "wb") if output_file:
file_to_save.write(data) with open(output_file, "wb") as file_to_save:
file_to_save.write(data)
else:
return data
except Exception as e: except Exception as e:
raise Exception(f"{__name__} error: {e}") raise Exception(f"{__name__} error: {e}")
@@ -22,6 +22,7 @@ class TTSProvider(TTSProviderBase):
self.api_url = "https://tts.tencentcloudapi.com" # 正确的API端点 self.api_url = "https://tts.tencentcloudapi.com" # 正确的API端点
self.region = config.get("region") self.region = config.get("region")
self.output_file = config.get("output_dir") self.output_file = config.get("output_dir")
self.audio_file_type = config.get("format", "wav")
def _get_auth_headers(self, request_body): def _get_auth_headers(self, request_body):
"""生成鉴权请求头""" """生成鉴权请求头"""
@@ -148,12 +149,14 @@ class TTSProvider(TTSProviderBase):
f"API返回错误: {error_info['Code']}: {error_info['Message']}" f"API返回错误: {error_info['Code']}: {error_info['Message']}"
) )
# 提取音频数据 # 解码Base64音频数据
audio_data = response_data["Response"].get("Audio") audio_bytes = base64.b64decode(response_data["Response"].get("Audio"))
if audio_data: if audio_bytes:
# 解码Base64音频数据并保存 if output_file:
with open(output_file, "wb") as f: with open(output_file, "wb") as f:
f.write(base64.b64decode(audio_data)) f.write(audio_bytes)
else:
return audio_bytes
else: else:
raise Exception(f"{__name__}: 没有返回音频数据: {response_data}") raise Exception(f"{__name__}: 没有返回音频数据: {response_data}")
else: else:
@@ -30,6 +30,7 @@ class TTSProvider(TTSProviderBase):
self.output_file = config.get("output_dir") self.output_file = config.get("output_dir")
self.pitch_factor = int(config.get("pitch_factor", 0)) self.pitch_factor = int(config.get("pitch_factor", 0))
self.format = config.get("format", "mp3") self.format = config.get("format", "mp3")
self.audio_file_type = config.get("format", "mp3")
self.emotion = int(config.get("emotion", 1)) self.emotion = int(config.get("emotion", 1))
self.header = {"Content-Type": "application/json"} self.header = {"Content-Type": "application/json"}
@@ -73,9 +74,11 @@ class TTSProvider(TTSProviderBase):
) )
audio_content = requests.get(result) audio_content = requests.get(result)
with open(output_file, "wb") as f: if output_file:
f.write(audio_content.content) with open(output_file, "wb") as f:
return True f.write(audio_content.content)
else:
return audio_content.content
voice_path = resp_json.get("voice_path") voice_path = resp_json.get("voice_path")
des_path = output_file des_path = output_file
shutil.move(voice_path, des_path) shutil.move(voice_path, des_path)
+26
View File
@@ -31,3 +31,29 @@ def decode_opus_from_file(input_file):
# 计算总时长 # 计算总时长
total_duration = (total_frames * frame_duration_ms) / 1000.0 total_duration = (total_frames * frame_duration_ms) / 1000.0
return opus_datas, total_duration return opus_datas, total_duration
def decode_opus_from_bytes(input_bytes):
"""
从p3二进制数据中解码 Opus 数据,并返回一个 Opus 数据包的列表以及总时长。
"""
import io
opus_datas = []
total_frames = 0
sample_rate = 16000 # 文件采样率
frame_duration_ms = 60 # 帧时长
frame_size = int(sample_rate * frame_duration_ms / 1000)
f = io.BytesIO(input_bytes)
while True:
header = f.read(4)
if not header:
break
_, _, data_len = struct.unpack('>BBH', header)
opus_data = f.read(data_len)
if len(opus_data) != data_len:
raise ValueError(f"Data length({len(opus_data)}) mismatch({data_len}) in the bytes.")
opus_datas.append(opus_data)
total_frames += 1
total_duration = (total_frames * frame_duration_ms) / 1000.0
return opus_datas, total_duration
+46
View File
@@ -3,6 +3,9 @@ import socket
import subprocess import subprocess
import re import re
import os import os
import wave
from io import BytesIO
from core.utils import p3
import numpy as np import numpy as np
import requests import requests
import opuslib_next import opuslib_next
@@ -773,6 +776,22 @@ def audio_to_data(audio_file_path, is_opus=True):
return pcm_to_data(raw_data, is_opus), duration return pcm_to_data(raw_data, is_opus), duration
def audio_bytes_to_data(audio_bytes, file_type, is_opus=True):
"""
直接用音频二进制数据转为opus/pcm数据,支持wav、mp3、p3
"""
if file_type == "p3":
# 直接用p3解码
return p3.decode_opus_from_bytes(audio_bytes)
else:
# 其他格式用pydub
audio = AudioSegment.from_file(BytesIO(audio_bytes), format=file_type, parameters=["-nostdin"])
audio = audio.set_channels(1).set_frame_rate(16000).set_sample_width(2)
duration = len(audio) / 1000.0
raw_data = audio.raw_data
return pcm_to_data(raw_data, is_opus), duration
def pcm_to_data(raw_data, is_opus=True): def pcm_to_data(raw_data, is_opus=True):
# 初始化Opus编码器 # 初始化Opus编码器
encoder = opuslib_next.Encoder(16000, 1, opuslib_next.APPLICATION_AUDIO) encoder = opuslib_next.Encoder(16000, 1, opuslib_next.APPLICATION_AUDIO)
@@ -804,6 +823,33 @@ def pcm_to_data(raw_data, is_opus=True):
return datas return datas
def opus_datas_to_wav_bytes(opus_datas, sample_rate=16000, channels=1):
"""
将opus帧列表解码为wav字节流
"""
decoder = opuslib_next.Decoder(sample_rate, channels)
pcm_datas = []
frame_duration = 60 # ms
frame_size = int(sample_rate * frame_duration / 1000) # 960
for opus_frame in opus_datas:
# 解码为PCM(返回bytes,2字节/采样点)
pcm = decoder.decode(opus_frame, frame_size)
pcm_datas.append(pcm)
pcm_bytes = b''.join(pcm_datas)
# 写入wav字节流
wav_buffer = BytesIO()
with wave.open(wav_buffer, 'wb') as wf:
wf.setnchannels(channels)
wf.setsampwidth(2) # 16bit
wf.setframerate(sample_rate)
wf.writeframes(pcm_bytes)
return wav_buffer.getvalue()
def check_vad_update(before_config, new_config): def check_vad_update(before_config, new_config):
if ( if (
new_config.get("selected_module") is None new_config.get("selected_module") is None