x
This commit is contained in:
Generated
+1
-1
@@ -3,5 +3,5 @@
|
||||
<component name="Black">
|
||||
<option name="sdkName" value="E:\anaconda3\envs\zghs_shuiwu_api" />
|
||||
</component>
|
||||
<component name="ProjectRootManager" version="2" project-jdk-name="python" project-jdk-type="Python SDK" />
|
||||
<component name="ProjectRootManager" version="2" project-jdk-name="AIStreamTest" project-jdk-type="Python SDK" />
|
||||
</project>
|
||||
Generated
+1
-1
@@ -2,7 +2,7 @@
|
||||
<module type="PYTHON_MODULE" version="4">
|
||||
<component name="NewModuleRootManager">
|
||||
<content url="file://$MODULE_DIR$" />
|
||||
<orderEntry type="jdk" jdkName="python" jdkType="Python SDK" />
|
||||
<orderEntry type="jdk" jdkName="AIStreamTest" jdkType="Python SDK" />
|
||||
<orderEntry type="sourceFolder" forTests="false" />
|
||||
</component>
|
||||
</module>
|
||||
@@ -0,0 +1,223 @@
|
||||
from enum import IntEnum
|
||||
import struct
|
||||
import json
|
||||
import gzip
|
||||
from typing import Optional, Union, Dict, Any, List
|
||||
|
||||
|
||||
# -------------------------- 协议常量定义 --------------------------
|
||||
class ProtocolConst:
|
||||
PROTOCOL_VERSION = 0b0001 # 协议版本v1
|
||||
HEADER_SIZE = 8 # 头部固定字节数(字节1~字节8)
|
||||
MAX_BODY_SIZE = 1024 * 1024 * 10 # 最大包体大小(10MB,防止内存溢出)
|
||||
STRING_ENCODING = "utf-8" # 字符串默认编码(统一编码格式,避免乱码)
|
||||
|
||||
|
||||
# -------------------------- 枚举定义(业务约定) --------------------------
|
||||
class MessageType(IntEnum):
|
||||
"""消息类型(4位,字节1高4位)"""
|
||||
PING = 0b0000 # 心跳
|
||||
AUDIO_DATA = 0b0001 # 纯音频数据(pcm)
|
||||
TEXT_MESSAGE = 0b0010 # 纯文本消息
|
||||
AUDIO_TEXT_MIX = 0b0011 # 音频+文本混合数据
|
||||
CONTROL_CMD = 0b0100 # 控制指令(暂停/继续等)
|
||||
|
||||
class SerializationType(IntEnum):
|
||||
"""序列化方式(4位,字节2高4位)"""
|
||||
RAW = 0b0000 # 原始二进制(音频等)
|
||||
JSON = 0b0001 # JSON格式(文本/指令)
|
||||
STRING = 0b0010 # 直接字符串(纯文本,UTF-8编码,无JSON包装)🔴新增类型
|
||||
|
||||
class CompressionType(IntEnum):
|
||||
"""压缩方式(4位,字节2低4位)"""
|
||||
NONE = 0b0000 # 无压缩
|
||||
GZIP = 0b0001 # gzip压缩
|
||||
|
||||
class ControlCommand(IntEnum):
|
||||
"""控制指令类型(配合MessageType.CONTROL_CMD使用)"""
|
||||
HEARTBEAT = 0b0001 # 心跳响应
|
||||
PAUSE = 0b0010 # 暂停
|
||||
RESUME = 0b0011 # 继续
|
||||
STOP = 0b0100 # 停止
|
||||
|
||||
# -------------------------- 协议工具类 --------------------------
|
||||
class ProtocolCodec:
|
||||
@staticmethod
|
||||
def pack(
|
||||
msg_type: MessageType,
|
||||
body: Union[bytes, str, Dict[str, Any], List[Any]],
|
||||
serialization: Optional[SerializationType] = None,
|
||||
compression: CompressionType = CompressionType.NONE
|
||||
) -> bytes:
|
||||
"""
|
||||
封装协议包(头部 + 包体)
|
||||
:param msg_type: 消息类型
|
||||
:param body: 包体数据(bytes/str/dict/list)
|
||||
:param serialization: 序列化方式(None时自动推导)
|
||||
:param compression: 压缩方式
|
||||
:return: 完整协议包(bytes)
|
||||
"""
|
||||
# 1. 自动推导序列化方式(优化:纯文本消息自动用STRING,复杂结构用JSON)🔴修改推导逻辑
|
||||
if serialization is None:
|
||||
if msg_type == MessageType.AUDIO_DATA:
|
||||
serialization = SerializationType.RAW
|
||||
elif msg_type == MessageType.TEXT_MESSAGE: # 纯文本消息→STRING
|
||||
serialization = SerializationType.STRING
|
||||
elif msg_type in (MessageType.CONTROL_CMD, MessageType.AUDIO_TEXT_MIX): # 复杂结构→JSON
|
||||
serialization = SerializationType.JSON
|
||||
else:
|
||||
raise ValueError(f"不支持的消息类型:{msg_type}")
|
||||
|
||||
# 2. 序列化包体(新增STRING类型处理)🔴新增分支
|
||||
serialized_body: bytes
|
||||
if serialization == SerializationType.RAW:
|
||||
if not isinstance(body, bytes):
|
||||
raise TypeError("RAW序列化要求body必须是bytes类型")
|
||||
serialized_body = body
|
||||
elif serialization == SerializationType.STRING:
|
||||
if not isinstance(body, str):
|
||||
raise TypeError("STRING序列化要求body必须是str类型")
|
||||
serialized_body = body.encode(ProtocolConst.STRING_ENCODING) # 直接UTF-8编码,无JSON包装
|
||||
elif serialization == SerializationType.JSON:
|
||||
if isinstance(body, str):
|
||||
serialized_body = body.encode(ProtocolConst.STRING_ENCODING)
|
||||
elif isinstance(body, (dict, list)):
|
||||
serialized_body = json.dumps(body, ensure_ascii=False, separators=(',', ':')).encode(ProtocolConst.STRING_ENCODING)
|
||||
else:
|
||||
raise TypeError("JSON序列化要求body必须是str/dict/list类型")
|
||||
else:
|
||||
raise ValueError(f"不支持的序列化方式:{serialization}")
|
||||
|
||||
# 3. 压缩包体(不变)
|
||||
compressed_body: bytes
|
||||
if compression == CompressionType.GZIP:
|
||||
compressed_body = gzip.compress(serialized_body)
|
||||
elif compression == CompressionType.NONE:
|
||||
compressed_body = serialized_body
|
||||
else:
|
||||
raise ValueError(f"不支持的压缩方式:{compression}")
|
||||
|
||||
# 4. 校验包体大小(不变)
|
||||
body_len = len(compressed_body)
|
||||
if body_len > ProtocolConst.MAX_BODY_SIZE:
|
||||
raise OverflowError(f"包体过大({body_len}字节),最大支持{ProtocolConst.MAX_BODY_SIZE}字节")
|
||||
|
||||
# 5. 构造头部(不变)
|
||||
byte1 = (msg_type.value << 4) | 0x00 # 消息类型(4位) + 保留位1(4位)
|
||||
byte2 = (serialization.value << 4) | (compression.value & 0x0F) # 序列化(4位) + 压缩(4位)
|
||||
byte3_4 = struct.pack(">H", ProtocolConst.PROTOCOL_VERSION) # 协议版本(16位)
|
||||
byte5_8 = struct.pack(">I", body_len) # 包体长度(32位)
|
||||
header = bytes([byte1, byte2]) + byte3_4 + byte5_8
|
||||
assert len(header) == ProtocolConst.HEADER_SIZE, f"头部长度错误:实际{len(header)}字节,预期{ProtocolConst.HEADER_SIZE}字节"
|
||||
|
||||
return header + compressed_body
|
||||
|
||||
@staticmethod
|
||||
def unpack(packet: bytes) -> tuple[MessageType, SerializationType, CompressionType, Any]:
|
||||
"""
|
||||
解析协议包
|
||||
:param packet: 完整协议包(头部 + 包体)
|
||||
:return: (消息类型, 序列化方式, 压缩方式, 原始包体数据)
|
||||
"""
|
||||
# 1. 校验包长度(不变)
|
||||
if len(packet) < ProtocolConst.HEADER_SIZE:
|
||||
raise ValueError(f"包长度过短({len(packet)}字节),至少需要{ProtocolConst.HEADER_SIZE}字节头部")
|
||||
|
||||
# 2. 解析头部(不变)
|
||||
header = packet[:ProtocolConst.HEADER_SIZE]
|
||||
body = packet[ProtocolConst.HEADER_SIZE:]
|
||||
byte1 = header[0]
|
||||
byte2 = header[1]
|
||||
|
||||
msg_type = MessageType((byte1 >> 4) & 0x0F)
|
||||
serialization = SerializationType((byte2 >> 4) & 0x0F)
|
||||
compression = CompressionType(byte2 & 0x0F)
|
||||
version = struct.unpack(">H", header[2:4])[0]
|
||||
body_len = struct.unpack(">I", header[4:8])[0]
|
||||
|
||||
if version != ProtocolConst.PROTOCOL_VERSION:
|
||||
raise ValueError(f"协议版本不匹配:收到v{version},支持v{ProtocolConst.PROTOCOL_VERSION}")
|
||||
if len(body) != body_len:
|
||||
raise ValueError(f"包体长度不匹配:头部声明{body_len}字节,实际{len(body)}字节")
|
||||
|
||||
# 3. 解压包体(不变)
|
||||
decompressed_body: bytes
|
||||
if compression == CompressionType.GZIP:
|
||||
try:
|
||||
decompressed_body = gzip.decompress(body)
|
||||
except Exception as e:
|
||||
raise ValueError(f"GZIP解压失败:{str(e)}")
|
||||
elif compression == CompressionType.NONE:
|
||||
decompressed_body = body
|
||||
else:
|
||||
raise ValueError(f"不支持的压缩方式:{compression}")
|
||||
|
||||
# 4. 反序列化包体(新增STRING类型处理)🔴新增分支
|
||||
original_body: Any
|
||||
if serialization == SerializationType.RAW:
|
||||
original_body = decompressed_body
|
||||
elif serialization == SerializationType.STRING:
|
||||
try:
|
||||
original_body = decompressed_body.decode(ProtocolConst.STRING_ENCODING) # 直接UTF-8解码,无JSON解析
|
||||
except UnicodeDecodeError:
|
||||
raise ValueError(f"STRING反序列化失败:{ProtocolConst.STRING_ENCODING}解码错误")
|
||||
elif serialization == SerializationType.JSON:
|
||||
try:
|
||||
original_body = json.loads(decompressed_body.decode(ProtocolConst.STRING_ENCODING))
|
||||
except UnicodeDecodeError:
|
||||
raise ValueError(f"JSON反序列化失败:{ProtocolConst.STRING_ENCODING}解码错误")
|
||||
except json.JSONDecodeError:
|
||||
raise ValueError("JSON反序列化失败:格式错误")
|
||||
else:
|
||||
raise ValueError(f"不支持的序列化方式:{serialization}")
|
||||
|
||||
return msg_type, serialization, compression, original_body
|
||||
|
||||
|
||||
# -------------------------- 使用示例(新增纯字符串消息测试) --------------------------
|
||||
if __name__ == "__main__":
|
||||
# 示例1:纯文本消息(自动用STRING序列化,无JSON包装)🔴测试新增类型
|
||||
text_body = "你好,这是纯字符串消息(无JSON)"
|
||||
text_packet = ProtocolCodec.pack(
|
||||
msg_type=MessageType.TEXT_MESSAGE,
|
||||
body=text_body,
|
||||
compression=CompressionType.NONE
|
||||
)
|
||||
print(f"纯文本消息包长度:{len(text_packet)}字节")
|
||||
msg_type1, ser1, comp1, body1 = ProtocolCodec.unpack(text_packet)
|
||||
print(f"解析结果:类型={msg_type1.name},序列化={ser1.name},压缩={comp1.name},内容={body1}\n")
|
||||
|
||||
# 示例2:控制指令(JSON序列化,复杂结构)
|
||||
control_body = {"cmd": ControlCommand.HEARTBEAT.value, "timestamp": 1699999999}
|
||||
control_packet = ProtocolCodec.pack(
|
||||
msg_type=MessageType.CONTROL_CMD,
|
||||
body=control_body,
|
||||
compression=CompressionType.NONE
|
||||
)
|
||||
print(f"控制指令包长度:{len(control_packet)}字节")
|
||||
msg_type2, ser2, comp2, body2 = ProtocolCodec.unpack(control_packet)
|
||||
print(f"解析结果:类型={msg_type2.name},序列化={ser2.name},压缩={comp2.name},内容={body2}\n")
|
||||
|
||||
# 示例3:音频数据(RAW序列化)
|
||||
audio_body = b"\x00\x01\x02\x03\x04\x05" * 100 # 模拟PCM数据
|
||||
audio_packet = ProtocolCodec.pack(
|
||||
msg_type=MessageType.AUDIO_DATA,
|
||||
body=audio_body,
|
||||
compression=CompressionType.NONE
|
||||
)
|
||||
print(f"音频数据包长度:{len(audio_packet)}字节")
|
||||
msg_type3, ser3, comp3, body3 = ProtocolCodec.unpack(audio_packet)
|
||||
print(f"解析结果:类型={msg_type3.name},序列化={ser3.name},压缩={comp3.name},数据长度={len(body3)}字节\n")
|
||||
|
||||
# 示例4:混合数据(JSON序列化)
|
||||
mix_body = {
|
||||
"audio_data": list(audio_body[:10]),
|
||||
"text_content": "混合数据仍用JSON"
|
||||
}
|
||||
mix_packet = ProtocolCodec.pack(
|
||||
msg_type=MessageType.AUDIO_TEXT_MIX,
|
||||
body=mix_body
|
||||
)
|
||||
print(f"混合数据包长度:{len(mix_packet)}字节")
|
||||
msg_type4, ser4, comp4, body4 = ProtocolCodec.unpack(mix_packet)
|
||||
print(f"解析结果:类型={msg_type4.name},序列化={ser4.name},压缩={comp4.name},内容={body4}")
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
+34
-32
@@ -66,57 +66,52 @@ async def frontend_websocket_handler(websocket: WebSocket):
|
||||
# ====================== 修复 TTS 核心逻辑 ======================
|
||||
# 1. 创建 TTS 客户端
|
||||
# 1. 创建TTS管理器
|
||||
tts_manager = TTSManager(ws_url="ws://10.10.10.202:50000/ws/tts")
|
||||
|
||||
# tts_manager = TTSManager(ws_url="ws://10.10.10.202:50000/ws/tts")
|
||||
tts_manager = TTSManager()
|
||||
# 2. 设置结果回调函数(接收完整结果)
|
||||
def handle_tts_result(req_id: str, result: Dict[str, Any]):
|
||||
"""处理TTS结果回调"""
|
||||
print('处理TTS结果回调 ',result)
|
||||
# print('处理TTS结果回调 ',result)
|
||||
status = result.get("status")
|
||||
|
||||
if status == "completed":
|
||||
audio_data = result.get("audio_data")
|
||||
sample_rate = result.get("sample_rate")
|
||||
|
||||
if audio_data is not None and len(audio_data) > 0:
|
||||
# 转换为PCM数据
|
||||
pcm_data = (audio_data.astype(np.float32) * 32767).astype(np.int16)
|
||||
pcm_bytes = pcm_data.tobytes()
|
||||
# result_queue.put_nowait({"type": 3, "data": pcm_bytes})
|
||||
result_queue.put_nowait(pcm_bytes)
|
||||
print(f"插入时候队列当前大小xxx: {result_queue.qsize()}") # 排查队列是否有数据
|
||||
result_queue.put_nowait(audio_data)
|
||||
#
|
||||
# if audio_data is not None and len(audio_data) > 0:
|
||||
# # 转换为PCM数据
|
||||
# pcm_data = (audio_data.astype(np.float32) * 32767).astype(np.int16)
|
||||
# pcm_bytes = pcm_data.tobytes()
|
||||
# # result_queue.put_nowait({"type": 3, "data": pcm_bytes})
|
||||
# result_queue.put_nowait(pcm_bytes)
|
||||
# print(f"插入时候队列当前大小xxx: {result_queue.qsize()}") # 排查队列是否有数据
|
||||
# 保存为PCM文件
|
||||
|
||||
|
||||
|
||||
elif status == "error":
|
||||
error_msg = result.get("message")
|
||||
print(f"❌ TTS处理失败 [{req_id[:8]}]: {error_msg}")
|
||||
# elif status == "error":
|
||||
# error_msg = result.get("message")
|
||||
# print(f"❌ TTS处理失败 [{req_id[:8]}]: {error_msg}")
|
||||
|
||||
tts_manager.set_result_callback(handle_tts_result)
|
||||
|
||||
# 3. 设置是否播放(可选,默认True)
|
||||
tts_manager.set_playback_enabled(True) # 设置为False则不播放
|
||||
tts_manager.set_playback_enabled(False) # 设置为False则不播放
|
||||
|
||||
# 4. 初始化连接
|
||||
await tts_manager.initialize()
|
||||
texts = [
|
||||
"你好,这是第一个排队的TTS请求。",
|
||||
"我是第二个请求,会等第一个处理完再执行。",
|
||||
"第三个请求,支持流式播放和队列管理。",
|
||||
"第四个请求,测试队列的自动消费功能。",
|
||||
"最后一个请求,处理完成后会自动结束。"
|
||||
]
|
||||
|
||||
req_ids = []
|
||||
for i, text in enumerate(texts):
|
||||
req_id = await tts_manager.synthesize(
|
||||
text,
|
||||
mode="预训练音色",
|
||||
sft_spk="中文女",
|
||||
speed=1.0
|
||||
)
|
||||
req_ids.append(req_id)
|
||||
# texts = [
|
||||
# "你好,这是第一个排队的TTS请求。你好,这是第一个排队的TTS请求。你好,这是第一个排队的TTS请好",
|
||||
# "我是第二个请求,会等第一个处理完再执行。",
|
||||
# "第三个请求,支持流式播放和队列管理。",
|
||||
# "第四个请求,测试队列的自动消费功能。",
|
||||
# "最后一个请求,处理完成后会自动结束。"
|
||||
# ]
|
||||
# req_ids = []
|
||||
# for i, text in enumerate(texts):
|
||||
# req_id = await tts_manager.synthesize(text)
|
||||
# req_ids.append(req_id)
|
||||
# ====================== ASR 结果回调 ======================
|
||||
async def asr_result_callback(result: dict):
|
||||
"""ASR 结果回调:转发前端 + 调用大模型"""
|
||||
@@ -151,6 +146,13 @@ async def frontend_websocket_handler(websocket: WebSocket):
|
||||
"""大模型流式回调(纯异步,无阻塞)"""
|
||||
if not chunk:
|
||||
return
|
||||
# req_id = await tts_manager.synthesize(
|
||||
# chunk,
|
||||
# mode="预训练音色",
|
||||
# sft_spk="中文女",
|
||||
# speed=1.0
|
||||
# )
|
||||
req_id = await tts_manager.synthesize(chunk)
|
||||
# 1. 异步插入队列(替代put_nowait,避免队列满时抛异常)
|
||||
# 1. 非阻塞插入(队列满则丢弃,优先保证实时性)
|
||||
# result_queue.put_nowait(chunk)
|
||||
|
||||
@@ -79,6 +79,7 @@ class LLMClient:
|
||||
response.raise_for_status()
|
||||
# 实时迭代流式响应
|
||||
async for line in response.content.iter_chunked(LLM_CONFIG["stream_chunk_size"]):
|
||||
|
||||
if not line:
|
||||
continue
|
||||
line_data = line.decode("utf-8")
|
||||
@@ -95,13 +96,14 @@ class LLMClient:
|
||||
res_conversation_id = data["conversation_id"]
|
||||
# 提取内容
|
||||
chunk = data.get("content", data.get("answer", data.get("message", "")))
|
||||
# print('大模型返回的', chunk)
|
||||
if chunk:
|
||||
full_response += chunk
|
||||
if stream_callback:
|
||||
await stream_callback(chunk, res_conversation_id, False)
|
||||
await asyncio.sleep(0) # 让出调度权
|
||||
except json.JSONDecodeError as e:
|
||||
print(f"解析流式数据失败: {e}")
|
||||
# print(f"大模型解析流式数据失败: {e}")
|
||||
continue
|
||||
else:
|
||||
# 异步非流式请求
|
||||
|
||||
@@ -0,0 +1,41 @@
|
||||
from .protocols import (
|
||||
CompressionBits,
|
||||
EventType,
|
||||
HeaderSizeBits,
|
||||
Message,
|
||||
MsgType,
|
||||
MsgTypeFlagBits,
|
||||
SerializationBits,
|
||||
VersionBits,
|
||||
audio_only_client,
|
||||
cancel_session,
|
||||
finish_connection,
|
||||
finish_session,
|
||||
full_client_request,
|
||||
receive_message,
|
||||
start_connection,
|
||||
start_session,
|
||||
task_request,
|
||||
wait_for_event,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"CompressionBits",
|
||||
"EventType",
|
||||
"HeaderSizeBits",
|
||||
"Message",
|
||||
"MsgType",
|
||||
"MsgTypeFlagBits",
|
||||
"SerializationBits",
|
||||
"VersionBits",
|
||||
"audio_only_client",
|
||||
"cancel_session",
|
||||
"finish_connection",
|
||||
"finish_session",
|
||||
"full_client_request",
|
||||
"receive_message",
|
||||
"start_connection",
|
||||
"start_session",
|
||||
"task_request",
|
||||
"wait_for_event",
|
||||
]
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,543 @@
|
||||
import io
|
||||
import logging
|
||||
import struct
|
||||
from dataclasses import dataclass
|
||||
from enum import IntEnum
|
||||
from typing import Callable, List
|
||||
|
||||
import websockets
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class MsgType(IntEnum):
|
||||
"""Message type enumeration"""
|
||||
|
||||
Invalid = 0
|
||||
FullClientRequest = 0b1
|
||||
AudioOnlyClient = 0b10
|
||||
FullServerResponse = 0b1001
|
||||
AudioOnlyServer = 0b1011
|
||||
FrontEndResultServer = 0b1100
|
||||
Error = 0b1111
|
||||
|
||||
# Alias
|
||||
ServerACK = AudioOnlyServer
|
||||
|
||||
def __str__(self) -> str:
|
||||
return self.name if self.name else f"MsgType({self.value})"
|
||||
|
||||
|
||||
class MsgTypeFlagBits(IntEnum):
|
||||
"""Message type flag bits"""
|
||||
|
||||
NoSeq = 0 # Non-terminal packet with no sequence
|
||||
PositiveSeq = 0b1 # Non-terminal packet with sequence > 0
|
||||
LastNoSeq = 0b10 # Last packet with no sequence
|
||||
NegativeSeq = 0b11 # Last packet with sequence < 0
|
||||
WithEvent = 0b100 # Payload contains event number (int32)
|
||||
|
||||
|
||||
class VersionBits(IntEnum):
|
||||
"""Version bits"""
|
||||
|
||||
Version1 = 1
|
||||
Version2 = 2
|
||||
Version3 = 3
|
||||
Version4 = 4
|
||||
|
||||
|
||||
class HeaderSizeBits(IntEnum):
|
||||
"""Header size bits"""
|
||||
|
||||
HeaderSize4 = 1
|
||||
HeaderSize8 = 2
|
||||
HeaderSize12 = 3
|
||||
HeaderSize16 = 4
|
||||
|
||||
|
||||
class SerializationBits(IntEnum):
|
||||
"""Serialization method bits"""
|
||||
|
||||
Raw = 0
|
||||
JSON = 0b1
|
||||
Thrift = 0b11
|
||||
Custom = 0b1111
|
||||
|
||||
|
||||
class CompressionBits(IntEnum):
|
||||
"""Compression method bits"""
|
||||
|
||||
None_ = 0
|
||||
Gzip = 0b1
|
||||
Custom = 0b1111
|
||||
|
||||
|
||||
class EventType(IntEnum):
|
||||
"""Event type enumeration"""
|
||||
|
||||
None_ = 0 # Default event
|
||||
|
||||
# 1 ~ 49 Upstream Connection events
|
||||
StartConnection = 1
|
||||
StartTask = 1 # Alias of StartConnection
|
||||
FinishConnection = 2
|
||||
FinishTask = 2 # Alias of FinishConnection
|
||||
|
||||
# 50 ~ 99 Downstream Connection events
|
||||
ConnectionStarted = 50 # Connection established successfully
|
||||
TaskStarted = 50 # Alias of ConnectionStarted
|
||||
ConnectionFailed = 51 # Connection failed (possibly due to authentication failure)
|
||||
TaskFailed = 51 # Alias of ConnectionFailed
|
||||
ConnectionFinished = 52 # Connection ended
|
||||
TaskFinished = 52 # Alias of ConnectionFinished
|
||||
|
||||
# 100 ~ 149 Upstream Session events
|
||||
StartSession = 100
|
||||
CancelSession = 101
|
||||
FinishSession = 102
|
||||
|
||||
# 150 ~ 199 Downstream Session events
|
||||
SessionStarted = 150
|
||||
SessionCanceled = 151
|
||||
SessionFinished = 152
|
||||
SessionFailed = 153
|
||||
UsageResponse = 154 # Usage response
|
||||
ChargeData = 154 # Alias of UsageResponse
|
||||
|
||||
# 200 ~ 249 Upstream general events
|
||||
TaskRequest = 200
|
||||
UpdateConfig = 201
|
||||
|
||||
# 250 ~ 299 Downstream general events
|
||||
AudioMuted = 250
|
||||
|
||||
# 300 ~ 349 Upstream TTS events
|
||||
SayHello = 300
|
||||
|
||||
# 350 ~ 399 Downstream TTS events
|
||||
TTSSentenceStart = 350
|
||||
TTSSentenceEnd = 351
|
||||
TTSResponse = 352
|
||||
TTSEnded = 359
|
||||
PodcastRoundStart = 360
|
||||
PodcastRoundResponse = 361
|
||||
PodcastRoundEnd = 362
|
||||
|
||||
# 450 ~ 499 Downstream ASR events
|
||||
ASRInfo = 450
|
||||
ASRResponse = 451
|
||||
ASREnded = 459
|
||||
|
||||
# 500 ~ 549 Upstream dialogue events
|
||||
ChatTTSText = 500 # (Ground-Truth-Alignment) text for speech synthesis
|
||||
|
||||
# 550 ~ 599 Downstream dialogue events
|
||||
ChatResponse = 550
|
||||
ChatEnded = 559
|
||||
|
||||
# 650 ~ 699 Downstream dialogue events
|
||||
# Events for source (original) language subtitle
|
||||
SourceSubtitleStart = 650
|
||||
SourceSubtitleResponse = 651
|
||||
SourceSubtitleEnd = 652
|
||||
# Events for target (translation) language subtitle
|
||||
TranslationSubtitleStart = 653
|
||||
TranslationSubtitleResponse = 654
|
||||
TranslationSubtitleEnd = 655
|
||||
|
||||
def __str__(self) -> str:
|
||||
return self.name if self.name else f"EventType({self.value})"
|
||||
|
||||
|
||||
@dataclass
|
||||
class Message:
|
||||
"""Message object
|
||||
|
||||
Message format:
|
||||
0 1 2 3
|
||||
| 0 1 2 3 4 5 6 7 | 0 1 2 3 4 5 6 7 | 0 1 2 3 4 5 6 7 | 0 1 2 3 4 5 6 7 |
|
||||
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
|
||||
| Version | Header Size | Msg Type | Flags |
|
||||
| (4 bits) | (4 bits) | (4 bits) | (4 bits) |
|
||||
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
|
||||
| Serialization | Compression | Reserved |
|
||||
| (4 bits) | (4 bits) | (8 bits) |
|
||||
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
|
||||
| |
|
||||
| Optional Header Extensions |
|
||||
| (if Header Size > 1) |
|
||||
| |
|
||||
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
|
||||
| |
|
||||
| Payload |
|
||||
| (variable length) |
|
||||
| |
|
||||
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
|
||||
"""
|
||||
|
||||
version: VersionBits = VersionBits.Version1
|
||||
header_size: HeaderSizeBits = HeaderSizeBits.HeaderSize4
|
||||
type: MsgType = MsgType.Invalid
|
||||
flag: MsgTypeFlagBits = MsgTypeFlagBits.NoSeq
|
||||
serialization: SerializationBits = SerializationBits.JSON
|
||||
compression: CompressionBits = CompressionBits.None_
|
||||
|
||||
event: EventType = EventType.None_
|
||||
session_id: str = ""
|
||||
connect_id: str = ""
|
||||
sequence: int = 0
|
||||
error_code: int = 0
|
||||
|
||||
payload: bytes = b""
|
||||
|
||||
@classmethod
|
||||
def from_bytes(cls, data: bytes) -> "Message":
|
||||
"""Create message object from bytes"""
|
||||
if len(data) < 3:
|
||||
raise ValueError(
|
||||
f"Data too short: expected at least 3 bytes, got {len(data)}"
|
||||
)
|
||||
|
||||
type_and_flag = data[1]
|
||||
msg_type = MsgType(type_and_flag >> 4)
|
||||
flag = MsgTypeFlagBits(type_and_flag & 0b00001111)
|
||||
|
||||
msg = cls(type=msg_type, flag=flag)
|
||||
msg.unmarshal(data)
|
||||
return msg
|
||||
|
||||
def marshal(self) -> bytes:
|
||||
"""Serialize message to bytes"""
|
||||
buffer = io.BytesIO()
|
||||
|
||||
# Write header
|
||||
header = [
|
||||
(self.version << 4) | self.header_size,
|
||||
(self.type << 4) | self.flag,
|
||||
(self.serialization << 4) | self.compression,
|
||||
]
|
||||
|
||||
header_size = 4 * self.header_size
|
||||
if padding := header_size - len(header):
|
||||
header.extend([0] * padding)
|
||||
|
||||
buffer.write(bytes(header))
|
||||
|
||||
# Write other fields
|
||||
writers = self._get_writers()
|
||||
for writer in writers:
|
||||
writer(buffer)
|
||||
|
||||
return buffer.getvalue()
|
||||
|
||||
def unmarshal(self, data: bytes) -> None:
|
||||
"""Deserialize message from bytes"""
|
||||
buffer = io.BytesIO(data)
|
||||
|
||||
# Read version and header size
|
||||
version_and_header_size = buffer.read(1)[0]
|
||||
self.version = VersionBits(version_and_header_size >> 4)
|
||||
self.header_size = HeaderSizeBits(version_and_header_size & 0b00001111)
|
||||
|
||||
# Skip second byte
|
||||
buffer.read(1)
|
||||
|
||||
# Read serialization and compression methods
|
||||
serialization_compression = buffer.read(1)[0]
|
||||
self.serialization = SerializationBits(serialization_compression >> 4)
|
||||
self.compression = CompressionBits(serialization_compression & 0b00001111)
|
||||
|
||||
# Skip header padding
|
||||
header_size = 4 * self.header_size
|
||||
read_size = 3
|
||||
if padding_size := header_size - read_size:
|
||||
buffer.read(padding_size)
|
||||
|
||||
# Read other fields
|
||||
readers = self._get_readers()
|
||||
for reader in readers:
|
||||
reader(buffer)
|
||||
|
||||
# Check for remaining data
|
||||
remaining = buffer.read()
|
||||
if remaining:
|
||||
raise ValueError(f"Unexpected data after message: {remaining}")
|
||||
|
||||
def _get_writers(self) -> List[Callable[[io.BytesIO], None]]:
|
||||
"""Get list of writer functions"""
|
||||
writers = []
|
||||
|
||||
if self.flag == MsgTypeFlagBits.WithEvent:
|
||||
writers.extend([self._write_event, self._write_session_id])
|
||||
|
||||
if self.type in [
|
||||
MsgType.FullClientRequest,
|
||||
MsgType.FullServerResponse,
|
||||
MsgType.FrontEndResultServer,
|
||||
MsgType.AudioOnlyClient,
|
||||
MsgType.AudioOnlyServer,
|
||||
]:
|
||||
if self.flag in [MsgTypeFlagBits.PositiveSeq, MsgTypeFlagBits.NegativeSeq]:
|
||||
writers.append(self._write_sequence)
|
||||
elif self.type == MsgType.Error:
|
||||
writers.append(self._write_error_code)
|
||||
else:
|
||||
raise ValueError(f"Unsupported message type: {self.type}")
|
||||
|
||||
writers.append(self._write_payload)
|
||||
return writers
|
||||
|
||||
def _get_readers(self) -> List[Callable[[io.BytesIO], None]]:
|
||||
"""Get list of reader functions"""
|
||||
readers = []
|
||||
|
||||
if self.type in [
|
||||
MsgType.FullClientRequest,
|
||||
MsgType.FullServerResponse,
|
||||
MsgType.FrontEndResultServer,
|
||||
MsgType.AudioOnlyClient,
|
||||
MsgType.AudioOnlyServer,
|
||||
]:
|
||||
if self.flag in [MsgTypeFlagBits.PositiveSeq, MsgTypeFlagBits.NegativeSeq]:
|
||||
readers.append(self._read_sequence)
|
||||
elif self.type == MsgType.Error:
|
||||
readers.append(self._read_error_code)
|
||||
else:
|
||||
raise ValueError(f"Unsupported message type: {self.type}")
|
||||
|
||||
if self.flag == MsgTypeFlagBits.WithEvent:
|
||||
readers.extend(
|
||||
[self._read_event, self._read_session_id, self._read_connect_id]
|
||||
)
|
||||
|
||||
readers.append(self._read_payload)
|
||||
return readers
|
||||
|
||||
def _write_event(self, buffer: io.BytesIO) -> None:
|
||||
"""Write event"""
|
||||
buffer.write(struct.pack(">i", self.event))
|
||||
|
||||
def _write_session_id(self, buffer: io.BytesIO) -> None:
|
||||
"""Write session ID"""
|
||||
if self.event in [
|
||||
EventType.StartConnection,
|
||||
EventType.FinishConnection,
|
||||
EventType.ConnectionStarted,
|
||||
EventType.ConnectionFailed,
|
||||
]:
|
||||
return
|
||||
|
||||
session_id_bytes = self.session_id.encode("utf-8")
|
||||
size = len(session_id_bytes)
|
||||
if size > 0xFFFFFFFF:
|
||||
raise ValueError(f"Session ID size ({size}) exceeds max(uint32)")
|
||||
|
||||
buffer.write(struct.pack(">I", size))
|
||||
if size > 0:
|
||||
buffer.write(session_id_bytes)
|
||||
|
||||
def _write_sequence(self, buffer: io.BytesIO) -> None:
|
||||
"""Write sequence number"""
|
||||
buffer.write(struct.pack(">i", self.sequence))
|
||||
|
||||
def _write_error_code(self, buffer: io.BytesIO) -> None:
|
||||
"""Write error code"""
|
||||
buffer.write(struct.pack(">I", self.error_code))
|
||||
|
||||
def _write_payload(self, buffer: io.BytesIO) -> None:
|
||||
"""Write payload"""
|
||||
size = len(self.payload)
|
||||
if size > 0xFFFFFFFF:
|
||||
raise ValueError(f"Payload size ({size}) exceeds max(uint32)")
|
||||
|
||||
buffer.write(struct.pack(">I", size))
|
||||
buffer.write(self.payload)
|
||||
|
||||
def _read_event(self, buffer: io.BytesIO) -> None:
|
||||
"""Read event"""
|
||||
event_bytes = buffer.read(4)
|
||||
if event_bytes:
|
||||
self.event = EventType(struct.unpack(">i", event_bytes)[0])
|
||||
|
||||
def _read_session_id(self, buffer: io.BytesIO) -> None:
|
||||
"""Read session ID"""
|
||||
if self.event in [
|
||||
EventType.StartConnection,
|
||||
EventType.FinishConnection,
|
||||
EventType.ConnectionStarted,
|
||||
EventType.ConnectionFailed,
|
||||
EventType.ConnectionFinished,
|
||||
]:
|
||||
return
|
||||
|
||||
size_bytes = buffer.read(4)
|
||||
if size_bytes:
|
||||
size = struct.unpack(">I", size_bytes)[0]
|
||||
if size > 0:
|
||||
session_id_bytes = buffer.read(size)
|
||||
if len(session_id_bytes) == size:
|
||||
self.session_id = session_id_bytes.decode("utf-8")
|
||||
|
||||
def _read_connect_id(self, buffer: io.BytesIO) -> None:
|
||||
"""Read connection ID"""
|
||||
if self.event in [
|
||||
EventType.ConnectionStarted,
|
||||
EventType.ConnectionFailed,
|
||||
EventType.ConnectionFinished,
|
||||
]:
|
||||
size_bytes = buffer.read(4)
|
||||
if size_bytes:
|
||||
size = struct.unpack(">I", size_bytes)[0]
|
||||
if size > 0:
|
||||
self.connect_id = buffer.read(size).decode("utf-8")
|
||||
|
||||
def _read_sequence(self, buffer: io.BytesIO) -> None:
|
||||
"""Read sequence number"""
|
||||
sequence_bytes = buffer.read(4)
|
||||
if sequence_bytes:
|
||||
self.sequence = struct.unpack(">i", sequence_bytes)[0]
|
||||
|
||||
def _read_error_code(self, buffer: io.BytesIO) -> None:
|
||||
"""Read error code"""
|
||||
error_code_bytes = buffer.read(4)
|
||||
if error_code_bytes:
|
||||
self.error_code = struct.unpack(">I", error_code_bytes)[0]
|
||||
|
||||
def _read_payload(self, buffer: io.BytesIO) -> None:
|
||||
"""Read payload"""
|
||||
size_bytes = buffer.read(4)
|
||||
if size_bytes:
|
||||
size = struct.unpack(">I", size_bytes)[0]
|
||||
if size > 0:
|
||||
self.payload = buffer.read(size)
|
||||
|
||||
def __str__(self) -> str:
|
||||
"""String representation"""
|
||||
if self.type in [MsgType.AudioOnlyServer, MsgType.AudioOnlyClient]:
|
||||
if self.flag in [MsgTypeFlagBits.PositiveSeq, MsgTypeFlagBits.NegativeSeq]:
|
||||
return f"MsgType: {self.type}, EventType:{self.event}, Sequence: {self.sequence}, PayloadSize: {len(self.payload)}"
|
||||
return f"MsgType: {self.type}, EventType:{self.event}, PayloadSize: {len(self.payload)}"
|
||||
elif self.type == MsgType.Error:
|
||||
return f"MsgType: {self.type}, EventType:{self.event}, ErrorCode: {self.error_code}, Payload: {self.payload.decode('utf-8', 'ignore')}"
|
||||
else:
|
||||
if self.flag in [MsgTypeFlagBits.PositiveSeq, MsgTypeFlagBits.NegativeSeq]:
|
||||
return f"MsgType: {self.type}, EventType:{self.event}, Sequence: {self.sequence}, Payload: {self.payload.decode('utf-8', 'ignore')}"
|
||||
return f"MsgType: {self.type}, EventType:{self.event}, Payload: {self.payload.decode('utf-8', 'ignore')}"
|
||||
|
||||
|
||||
async def receive_message(websocket: websockets.WebSocketClientProtocol) -> Message:
|
||||
"""Receive message from websocket"""
|
||||
try:
|
||||
data = await websocket.recv()
|
||||
if isinstance(data, str):
|
||||
raise ValueError(f"Unexpected text message: {data}")
|
||||
elif isinstance(data, bytes):
|
||||
msg = Message.from_bytes(data)
|
||||
logger.info(f"Received: {msg}")
|
||||
return msg
|
||||
else:
|
||||
raise ValueError(f"Unexpected message type: {type(data)}")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to receive message: {e}")
|
||||
raise
|
||||
|
||||
|
||||
async def wait_for_event(
|
||||
websocket: websockets.WebSocketClientProtocol,
|
||||
msg_type: MsgType,
|
||||
event_type: EventType,
|
||||
) -> Message:
|
||||
"""Wait for specific event"""
|
||||
while True:
|
||||
msg = await receive_message(websocket)
|
||||
if msg.type != msg_type or msg.event != event_type:
|
||||
raise ValueError(f"Unexpected message: {msg}")
|
||||
if msg.type == msg_type and msg.event == event_type:
|
||||
return msg
|
||||
|
||||
|
||||
async def full_client_request(
|
||||
websocket: websockets.WebSocketClientProtocol, payload: bytes
|
||||
) -> None:
|
||||
"""Send full client message"""
|
||||
msg = Message(type=MsgType.FullClientRequest, flag=MsgTypeFlagBits.NoSeq)
|
||||
msg.payload = payload
|
||||
logger.info(f"Sending: {msg}")
|
||||
await websocket.send(msg.marshal())
|
||||
|
||||
|
||||
async def audio_only_client(
|
||||
websocket: websockets.WebSocketClientProtocol, payload: bytes, flag: MsgTypeFlagBits
|
||||
) -> None:
|
||||
"""Send audio-only client message"""
|
||||
msg = Message(type=MsgType.AudioOnlyClient, flag=flag)
|
||||
msg.payload = payload
|
||||
logger.info(f"Sending: {msg}")
|
||||
await websocket.send(msg.marshal())
|
||||
|
||||
|
||||
async def start_connection(websocket: websockets.WebSocketClientProtocol) -> None:
|
||||
"""Start connection"""
|
||||
msg = Message(type=MsgType.FullClientRequest, flag=MsgTypeFlagBits.WithEvent)
|
||||
msg.event = EventType.StartConnection
|
||||
msg.payload = b"{}"
|
||||
logger.info(f"Sending: {msg}")
|
||||
await websocket.send(msg.marshal())
|
||||
|
||||
|
||||
async def finish_connection(websocket: websockets.WebSocketClientProtocol) -> None:
|
||||
"""Finish connection"""
|
||||
msg = Message(type=MsgType.FullClientRequest, flag=MsgTypeFlagBits.WithEvent)
|
||||
msg.event = EventType.FinishConnection
|
||||
msg.payload = b"{}"
|
||||
logger.info(f"Sending: {msg}")
|
||||
await websocket.send(msg.marshal())
|
||||
|
||||
|
||||
async def start_session(
|
||||
websocket: websockets.WebSocketClientProtocol, payload: bytes, session_id: str
|
||||
) -> None:
|
||||
"""Start session"""
|
||||
msg = Message(type=MsgType.FullClientRequest, flag=MsgTypeFlagBits.WithEvent)
|
||||
msg.event = EventType.StartSession
|
||||
msg.session_id = session_id
|
||||
msg.payload = payload
|
||||
logger.info(f"Sending: {msg}")
|
||||
await websocket.send(msg.marshal())
|
||||
|
||||
|
||||
async def finish_session(
|
||||
websocket: websockets.WebSocketClientProtocol, session_id: str
|
||||
) -> None:
|
||||
"""Finish session"""
|
||||
msg = Message(type=MsgType.FullClientRequest, flag=MsgTypeFlagBits.WithEvent)
|
||||
msg.event = EventType.FinishSession
|
||||
msg.session_id = session_id
|
||||
msg.payload = b"{}"
|
||||
logger.info(f"Sending: {msg}")
|
||||
await websocket.send(msg.marshal())
|
||||
|
||||
|
||||
async def cancel_session(
|
||||
websocket: websockets.WebSocketClientProtocol, session_id: str
|
||||
) -> None:
|
||||
"""Cancel session"""
|
||||
msg = Message(type=MsgType.FullClientRequest, flag=MsgTypeFlagBits.WithEvent)
|
||||
msg.event = EventType.CancelSession
|
||||
msg.session_id = session_id
|
||||
msg.payload = b"{}"
|
||||
logger.info(f"Sending: {msg}")
|
||||
await websocket.send(msg.marshal())
|
||||
|
||||
|
||||
async def task_request(
|
||||
websocket: websockets.WebSocketClientProtocol, payload: bytes, session_id: str
|
||||
) -> None:
|
||||
"""Send task request"""
|
||||
msg = Message(type=MsgType.FullClientRequest, flag=MsgTypeFlagBits.WithEvent)
|
||||
msg.event = EventType.TaskRequest
|
||||
msg.session_id = session_id
|
||||
msg.payload = payload
|
||||
logger.info(f"Sending: {msg}")
|
||||
await websocket.send(msg.marshal())
|
||||
@@ -0,0 +1,551 @@
|
||||
import asyncio
|
||||
import json
|
||||
import websockets
|
||||
import numpy as np
|
||||
import sounddevice as sd
|
||||
from typing import Optional, Callable, Dict, Any, Coroutine, List
|
||||
from dataclasses import dataclass, field
|
||||
import uuid
|
||||
|
||||
|
||||
@dataclass
|
||||
class TTSRequest:
|
||||
"""TTS请求对象(带唯一标识)"""
|
||||
tts_text: str
|
||||
mode: str = "预训练音色"
|
||||
sft_spk: str = ""
|
||||
seed: int = field(default_factory=lambda: np.random.randint(1, 100000000))
|
||||
stream: bool = True
|
||||
speed: float = 1.0
|
||||
prompt_wav: str = ""
|
||||
instruct_text: str = ""
|
||||
request_id: str = field(default_factory=lambda: str(uuid.uuid4())) # 唯一请求ID
|
||||
|
||||
|
||||
class CosyVoiceTTSSocketClient:
|
||||
"""CosyVoice TTS WebSocket 客户端(异步/流式/带任务队列)"""
|
||||
|
||||
def __init__(self, ws_url: str = "ws://localhost:50000/ws/tts", max_queue_size: int = 100):
|
||||
"""
|
||||
初始化客户端
|
||||
:param ws_url: WebSocket 服务端地址
|
||||
:param max_queue_size: 最大队列长度(防止内存溢出)
|
||||
"""
|
||||
self.ws_url = ws_url
|
||||
self.websocket: Optional[websockets.WebSocketClientProtocol] = None
|
||||
self.is_connected = False
|
||||
self.is_processing = False # 是否正在处理请求
|
||||
|
||||
# 异步任务队列(FIFO)
|
||||
self.request_queue: asyncio.Queue[TTSRequest] = asyncio.Queue(maxsize=max_queue_size)
|
||||
|
||||
# 回调函数定义(所有回调都带request_id,方便关联请求)
|
||||
self.on_task_enqueue: Callable[[str], None] = lambda req_id: None # 任务入队回调
|
||||
self.on_start: Callable[[str, Dict[str, Any]], None] = lambda req_id, data: None # 合成开始回调
|
||||
self.on_audio_chunk: Callable[[str, np.ndarray], None] = lambda req_id, chunk: None # 音频块回调
|
||||
self.on_end: Callable[[str, Dict[str, Any]], None] = lambda req_id, data: None # 合成结束回调
|
||||
self.on_error: Callable[[str, str], None] = lambda req_id, msg: None # 错误回调
|
||||
self.on_queue_full: Callable[[str], None] = lambda req_id: None # 队列满回调
|
||||
|
||||
# 音频播放流
|
||||
self.play_stream: Optional[sd.OutputStream] = None
|
||||
self.current_req_id: Optional[str] = None
|
||||
|
||||
# 是否启用播放功能
|
||||
self.enable_playback: bool = True
|
||||
|
||||
# 外部回调函数,用于返回完整结果
|
||||
self.external_callback: Optional[Callable[[str, Dict[str, Any]], None]] = None
|
||||
|
||||
# 存储每个请求的音频数据
|
||||
self.audio_buffers: Dict[str, List[np.ndarray]] = {}
|
||||
|
||||
async def connect(self):
|
||||
"""建立 WebSocket 连接(初始化一次)"""
|
||||
if not self.is_connected:
|
||||
try:
|
||||
self.websocket = await websockets.connect(self.ws_url)
|
||||
self.is_connected = True
|
||||
print(f"成功连接到 TTS 服务端: {self.ws_url}")
|
||||
|
||||
# 启动队列消费协程(后台运行)
|
||||
asyncio.create_task(self._consume_queue())
|
||||
except Exception as e:
|
||||
raise ConnectionError(f"连接失败: {str(e)}")
|
||||
|
||||
async def disconnect(self):
|
||||
"""关闭 WebSocket 连接"""
|
||||
if self.is_connected and self.websocket:
|
||||
await self.websocket.close()
|
||||
self.is_connected = False
|
||||
self.websocket = None
|
||||
print("已断开与 TTS 服务端的连接")
|
||||
|
||||
if self.play_stream:
|
||||
self.play_stream.stop()
|
||||
self.play_stream.close()
|
||||
self.play_stream = None
|
||||
|
||||
def set_external_callback(self, callback: Callable[[str, Dict[str, Any]], None]):
|
||||
"""设置外部回调函数,用于返回完整结果"""
|
||||
self.external_callback = callback
|
||||
|
||||
def set_playback_enabled(self, enabled: bool):
|
||||
"""设置是否启用音频播放"""
|
||||
self.enable_playback = enabled
|
||||
|
||||
async def add_tts_request(self, tts_text: str, **kwargs) -> str:
|
||||
"""
|
||||
添加TTS请求到队列(异步非阻塞)
|
||||
:param tts_text: 要合成的文本
|
||||
:param kwargs: 其他TTS参数
|
||||
:return: 唯一请求ID
|
||||
"""
|
||||
# 创建请求对象
|
||||
request = TTSRequest(tts_text=tts_text, **kwargs)
|
||||
req_id = request.request_id
|
||||
|
||||
# 初始化音频缓冲区
|
||||
self.audio_buffers[req_id] = []
|
||||
|
||||
# 异步入队
|
||||
try:
|
||||
await self.request_queue.put(request)
|
||||
self.on_task_enqueue(req_id)
|
||||
print(f"请求 {req_id} 已加入队列,当前队列长度: {self.request_queue.qsize()}")
|
||||
return req_id
|
||||
except asyncio.QueueFull:
|
||||
self.on_queue_full(req_id)
|
||||
raise Exception(f"队列已满,请求 {req_id} 入队失败")
|
||||
|
||||
async def _consume_queue(self):
|
||||
"""消费队列(后台协程,自动处理排队请求)"""
|
||||
print("队列消费协程已启动")
|
||||
while True:
|
||||
try:
|
||||
# 等待队列中有请求(阻塞)
|
||||
request = await self.request_queue.get()
|
||||
req_id = request.request_id
|
||||
|
||||
# 标记为处理中
|
||||
self.is_processing = True
|
||||
print(f"开始处理请求 {req_id},剩余队列长度: {self.request_queue.qsize()}")
|
||||
|
||||
# 处理当前请求
|
||||
await self._process_single_request(request)
|
||||
|
||||
# 标记任务完成
|
||||
self.request_queue.task_done()
|
||||
self.is_processing = False
|
||||
|
||||
except Exception as e:
|
||||
print(f"队列消费异常: {str(e)}")
|
||||
self.is_processing = False
|
||||
# 短暂等待后继续消费,避免死循环
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
async def _process_single_request(self, request: TTSRequest):
|
||||
"""处理单个TTS请求"""
|
||||
req_id = request.request_id
|
||||
|
||||
# 参数校验
|
||||
if not request.tts_text:
|
||||
self.on_error(req_id, "合成文本不能为空")
|
||||
return
|
||||
|
||||
# 确保已连接
|
||||
if not self.is_connected:
|
||||
try:
|
||||
await self.connect()
|
||||
except Exception as e:
|
||||
self.on_error(req_id, f"连接失败: {str(e)}")
|
||||
return
|
||||
|
||||
# 构造请求数据
|
||||
request_data = {
|
||||
"tts_text": request.tts_text,
|
||||
"mode": request.mode,
|
||||
"sft_spk": request.sft_spk,
|
||||
"seed": request.seed,
|
||||
"stream": request.stream,
|
||||
"speed": request.speed,
|
||||
"prompt_wav": request.prompt_wav,
|
||||
"instruct_text": request.instruct_text
|
||||
}
|
||||
|
||||
try:
|
||||
# 发送请求
|
||||
await self.websocket.send(json.dumps(request_data))
|
||||
|
||||
# 处理返回的流式数据
|
||||
await self._handle_response(req_id)
|
||||
|
||||
except websockets.exceptions.ConnectionClosed:
|
||||
self.is_connected = False
|
||||
self.on_error(req_id, "连接已关闭")
|
||||
except Exception as e:
|
||||
self.on_error(req_id, f"处理请求失败: {str(e)}")
|
||||
|
||||
async def _handle_response(self, req_id: str):
|
||||
"""处理单个请求的服务端响应"""
|
||||
if not self.websocket:
|
||||
return
|
||||
|
||||
try:
|
||||
sample_rate = None
|
||||
while True:
|
||||
# 接收服务端消息(异步)
|
||||
response = await self.websocket.recv()
|
||||
data = json.loads(response)
|
||||
|
||||
# 根据状态分发到不同回调(都带req_id)
|
||||
status = data.get("status")
|
||||
if status == "start":
|
||||
# 合成开始 - 返回采样率等信息
|
||||
sample_rate = data.get("sample_rate")
|
||||
self.on_start(req_id, data)
|
||||
|
||||
elif status == "stream":
|
||||
# 流式音频块 - 转为 numpy 数组
|
||||
audio_chunk = np.array(data["audio_chunk"], dtype=np.float32)
|
||||
|
||||
# 保存到缓冲区
|
||||
self.audio_buffers[req_id].append(audio_chunk)
|
||||
|
||||
# 音频块回调
|
||||
self.on_audio_chunk(req_id, audio_chunk)
|
||||
|
||||
elif status == "end":
|
||||
# 合成结束
|
||||
self.on_end(req_id, data)
|
||||
|
||||
# 组装完整音频数据
|
||||
full_audio = np.concatenate(self.audio_buffers[req_id]) if self.audio_buffers[req_id] else np.array(
|
||||
[], dtype=np.float32)
|
||||
|
||||
# 准备返回给外部回调的数据
|
||||
result_data = {
|
||||
"status": "completed",
|
||||
"request_id": req_id,
|
||||
"sample_rate": sample_rate,
|
||||
"audio_data": full_audio,
|
||||
"audio_length": len(full_audio),
|
||||
"message": data.get("msg", "合成完成")
|
||||
}
|
||||
|
||||
# 调用外部回调(非阻塞)
|
||||
if self.external_callback:
|
||||
# 使用create_task确保回调不会阻塞主流程
|
||||
asyncio.create_task(self._call_external_callback(req_id, result_data))
|
||||
|
||||
# 清理缓冲区
|
||||
if req_id in self.audio_buffers:
|
||||
del self.audio_buffers[req_id]
|
||||
|
||||
break
|
||||
|
||||
elif status == "error":
|
||||
# 错误处理
|
||||
error_msg = data.get("msg", "未知错误")
|
||||
self.on_error(req_id, error_msg)
|
||||
|
||||
# 准备错误结果数据
|
||||
error_data = {
|
||||
"status": "error",
|
||||
"request_id": req_id,
|
||||
"message": error_msg
|
||||
}
|
||||
|
||||
# 调用外部回调
|
||||
if self.external_callback:
|
||||
asyncio.create_task(self._call_external_callback(req_id, error_data))
|
||||
|
||||
# 清理缓冲区
|
||||
if req_id in self.audio_buffers:
|
||||
del self.audio_buffers[req_id]
|
||||
|
||||
break
|
||||
|
||||
except Exception as e:
|
||||
error_msg = f"接收数据异常: {str(e)}"
|
||||
self.on_error(req_id, error_msg)
|
||||
|
||||
# 准备错误结果数据
|
||||
error_data = {
|
||||
"status": "error",
|
||||
"request_id": req_id,
|
||||
"message": error_msg
|
||||
}
|
||||
|
||||
# 调用外部回调
|
||||
if self.external_callback:
|
||||
asyncio.create_task(self._call_external_callback(req_id, error_data))
|
||||
|
||||
# 清理缓冲区
|
||||
if req_id in self.audio_buffers:
|
||||
del self.audio_buffers[req_id]
|
||||
|
||||
async def _call_external_callback(self, req_id: str, result_data: Dict[str, Any]):
|
||||
"""异步调用外部回调函数"""
|
||||
try:
|
||||
if self.external_callback:
|
||||
# 确保回调是异步的,不会阻塞
|
||||
if asyncio.iscoroutinefunction(self.external_callback):
|
||||
await self.external_callback(req_id, result_data)
|
||||
else:
|
||||
# 如果是同步函数,在线程池中执行
|
||||
await asyncio.get_event_loop().run_in_executor(
|
||||
None, self.external_callback, req_id, result_data
|
||||
)
|
||||
except Exception as e:
|
||||
print(f"外部回调执行异常: {str(e)}")
|
||||
|
||||
def _play_audio_chunk(self, chunk: np.ndarray):
|
||||
"""播放音频块(非阻塞)"""
|
||||
if not self.enable_playback or not self.play_stream:
|
||||
return
|
||||
|
||||
try:
|
||||
if chunk.size > 0:
|
||||
self.play_stream.write(chunk)
|
||||
except Exception as e:
|
||||
print(f"音频播放异常: {str(e)}")
|
||||
|
||||
async def wait_queue_empty(self):
|
||||
"""等待队列所有任务处理完成(阻塞)"""
|
||||
await self.request_queue.join()
|
||||
print("所有队列任务已处理完成")
|
||||
|
||||
|
||||
# ------------------------------
|
||||
# 使用示例
|
||||
# ------------------------------
|
||||
class TTSManager:
|
||||
"""TTS管理器 - 供外部代码调用"""
|
||||
|
||||
def __init__(self, ws_url: str = "ws://localhost:50000/ws/tts"):
|
||||
self.client = CosyVoiceTTSSocketClient(ws_url)
|
||||
self._setup_callbacks()
|
||||
|
||||
def _setup_callbacks(self):
|
||||
"""设置内部回调"""
|
||||
|
||||
# 合成开始回调
|
||||
def on_tts_start(req_id, data):
|
||||
print(f"\n🎤 开始合成 [{req_id[:8]}] - 采样率: {data.get('sample_rate')}")
|
||||
|
||||
# 初始化播放流(如果需要播放)
|
||||
if self.client.enable_playback:
|
||||
sample_rate = data.get("sample_rate", 24000)
|
||||
if self.client.play_stream:
|
||||
self.client.play_stream.stop()
|
||||
self.client.play_stream.close()
|
||||
|
||||
self.client.play_stream = sd.OutputStream(
|
||||
samplerate=sample_rate,
|
||||
channels=1,
|
||||
dtype=np.float32
|
||||
)
|
||||
self.client.play_stream.start()
|
||||
self.client.current_req_id = req_id
|
||||
|
||||
# 音频块回调
|
||||
def on_audio_chunk(req_id, chunk):
|
||||
if req_id == self.client.current_req_id:
|
||||
self.client._play_audio_chunk(chunk)
|
||||
|
||||
# 合成结束回调
|
||||
def on_tts_end(req_id, data):
|
||||
print(f"🏁 合成完成 [{req_id[:8]}] - {data.get('msg', '完成')}")
|
||||
|
||||
# 错误回调
|
||||
def on_tts_error(req_id, msg):
|
||||
print(f"❌ 合成失败 [{req_id[:8]}] - {msg}")
|
||||
|
||||
self.client.on_start = on_tts_start
|
||||
self.client.on_audio_chunk = on_audio_chunk
|
||||
self.client.on_end = on_tts_end
|
||||
self.client.on_error = on_tts_error
|
||||
|
||||
async def initialize(self):
|
||||
"""初始化连接"""
|
||||
await self.client.connect()
|
||||
|
||||
async def shutdown(self):
|
||||
"""关闭连接"""
|
||||
await self.client.disconnect()
|
||||
|
||||
def set_result_callback(self, callback: Callable[[str, Dict[str, Any]], None]):
|
||||
"""设置结果回调函数"""
|
||||
self.client.set_external_callback(callback)
|
||||
|
||||
def set_playback_enabled(self, enabled: bool):
|
||||
"""设置是否启用播放"""
|
||||
self.client.set_playback_enabled(enabled)
|
||||
|
||||
async def synthesize(self, text: str, **kwargs) -> str:
|
||||
"""
|
||||
异步合成文本(非阻塞)
|
||||
:param text: 要合成的文本
|
||||
:param kwargs: 其他参数(mode, sft_spk, speed等)
|
||||
:return: 请求ID
|
||||
"""
|
||||
return await self.client.add_tts_request(text, **kwargs)
|
||||
|
||||
async def wait_all_completed(self):
|
||||
"""等待所有任务完成"""
|
||||
await self.client.wait_queue_empty()
|
||||
|
||||
|
||||
# ------------------------------
|
||||
# 外部调用示例
|
||||
# ------------------------------
|
||||
async def external_usage_example():
|
||||
"""外部代码使用示例"""
|
||||
|
||||
# 1. 创建TTS管理器
|
||||
tts_manager = TTSManager(ws_url="ws://10.10.10.202:50000/ws/tts")
|
||||
|
||||
# 2. 设置结果回调函数(接收完整结果)
|
||||
def handle_tts_result1(req_id: str, result: Dict[str, Any]):
|
||||
"""处理TTS结果回调"""
|
||||
status = result.get("status")
|
||||
|
||||
if status == "completed":
|
||||
audio_data = result.get("audio_data")
|
||||
sample_rate = result.get("sample_rate")
|
||||
audio_length = result.get("audio_length")
|
||||
|
||||
print(f"✅ 收到TTS结果 [{req_id[:8]}]: 长度{audio_length}采样点, 采样率{sample_rate}Hz")
|
||||
|
||||
# 这里可以保存音频文件或进行其他处理
|
||||
# 注意:audio_data是完整的numpy数组
|
||||
|
||||
elif status == "error":
|
||||
error_msg = result.get("message")
|
||||
print(f"❌ TTS处理失败 [{req_id[:8]}]: {error_msg}")
|
||||
|
||||
def handle_tts_result(req_id: str, result: Dict[str, Any]):
|
||||
"""处理TTS结果回调 - 保存为PCM文件"""
|
||||
print('处理TTS结果回调', result)
|
||||
status = result.get("status")
|
||||
print('status')
|
||||
if status == "completed":
|
||||
audio_data = result.get("audio_data")
|
||||
sample_rate = result.get("sample_rate")
|
||||
|
||||
if audio_data is not None and len(audio_data) > 0:
|
||||
# 转换为PCM数据
|
||||
pcm_data = (audio_data.astype(np.float32) * 32767).astype(np.int16)
|
||||
pcm_bytes = pcm_data.tobytes()
|
||||
|
||||
# 保存为PCM文件
|
||||
filename = f"tts_output_{req_id[:8]}.pcm"
|
||||
with open(filename, "wb") as f:
|
||||
f.write(pcm_bytes)
|
||||
|
||||
print(f"💾 PCM文件已保存: {filename}, 大小: {len(pcm_bytes)}字节")
|
||||
|
||||
return filename
|
||||
|
||||
elif status == "error":
|
||||
error_msg = result.get("message")
|
||||
print(f"❌ TTS处理失败 [{req_id[:8]}]: {error_msg}")
|
||||
|
||||
return None
|
||||
tts_manager.set_result_callback(handle_tts_result)
|
||||
|
||||
# 3. 设置是否播放(可选,默认True)
|
||||
tts_manager.set_playback_enabled(True) # 设置为False则不播放
|
||||
|
||||
# 4. 初始化连接
|
||||
await tts_manager.initialize()
|
||||
|
||||
# 5. 异步合成多个文本(非阻塞)
|
||||
texts = [
|
||||
"你好,这是第一个排队的TTS请求。",
|
||||
"我是第二个请求,会等第一个处理完再执行。",
|
||||
"第三个请求,支持流式播放和队列管理。",
|
||||
"第四个请求,测试队列的自动消费功能。",
|
||||
"最后一个请求,处理完成后会自动结束。"
|
||||
]
|
||||
|
||||
req_ids = []
|
||||
for i, text in enumerate(texts):
|
||||
req_id = await tts_manager.synthesize(
|
||||
text,
|
||||
mode="预训练音色",
|
||||
sft_spk="中文女",
|
||||
speed=1.0
|
||||
)
|
||||
req_ids.append(req_id)
|
||||
print(f"已提交请求 {i + 1}: ID={req_id[:8]}")
|
||||
|
||||
# 可以立即继续其他操作,不需要等待
|
||||
await asyncio.sleep(0.5) # 模拟其他操作
|
||||
|
||||
# 6. 可以在这里做其他事情,TTS会在后台处理
|
||||
|
||||
# 7. 等待所有TTS任务完成(可选)
|
||||
await tts_manager.wait_all_completed()
|
||||
|
||||
# 8. 关闭连接
|
||||
await tts_manager.shutdown()
|
||||
|
||||
|
||||
async def advanced_usage_example():
|
||||
"""高级使用示例:动态控制播放和结果处理"""
|
||||
|
||||
tts_manager = TTSManager(ws_url="ws://10.10.10.202:50000/ws/tts")
|
||||
|
||||
# 自定义结果处理器
|
||||
class ResultProcessor:
|
||||
def __init__(self):
|
||||
self.results = {}
|
||||
|
||||
async def process_result(self, req_id: str, result: Dict[str, Any]):
|
||||
"""异步处理结果"""
|
||||
if result.get("status") == "completed":
|
||||
# 保存音频数据等处理
|
||||
self.results[req_id] = result
|
||||
print(f"处理完成: {req_id[:8]}, 音频长度: {result.get('audio_length')}")
|
||||
|
||||
# 这里可以添加自定义逻辑,如保存到文件、发送到网络等
|
||||
|
||||
processor = ResultProcessor()
|
||||
tts_manager.set_result_callback(processor.process_result)
|
||||
|
||||
# 禁用播放,只获取数据
|
||||
tts_manager.set_playback_enabled(False)
|
||||
|
||||
await tts_manager.initialize()
|
||||
|
||||
# 批量合成
|
||||
tasks = []
|
||||
for i in range(3):
|
||||
text = f"这是第{i+1}个测试文本,用于演示高级功能。"
|
||||
task = tts_manager.synthesize(text)
|
||||
tasks.append(task)
|
||||
|
||||
# 并行提交所有请求
|
||||
req_ids = await asyncio.gather(*tasks)
|
||||
|
||||
# 等待完成
|
||||
await tts_manager.wait_all_completed()
|
||||
await tts_manager.shutdown()
|
||||
|
||||
print(f"所有任务完成,共处理 {len(processor.results)} 个结果")
|
||||
|
||||
if __name__ == "__main__":
|
||||
# 运行示例
|
||||
try:
|
||||
# 基础使用示例
|
||||
asyncio.run(external_usage_example())
|
||||
|
||||
# 高级使用示例(取消注释运行)
|
||||
# asyncio.run(advanced_usage_example())
|
||||
|
||||
except KeyboardInterrupt:
|
||||
print("\n程序被用户中断")
|
||||
except Exception as e:
|
||||
print(f"程序异常: {str(e)}")
|
||||
+458
-315
@@ -3,38 +3,73 @@ import json
|
||||
import websockets
|
||||
import numpy as np
|
||||
import sounddevice as sd
|
||||
from typing import Optional, Callable, Dict, Any, Coroutine, List
|
||||
from typing import Optional, Callable, Dict, Any, List
|
||||
from dataclasses import dataclass, field
|
||||
import uuid
|
||||
import copy
|
||||
from protocols import (
|
||||
EventType,
|
||||
MsgType,
|
||||
finish_connection,
|
||||
finish_session,
|
||||
receive_message,
|
||||
start_connection,
|
||||
start_session,
|
||||
task_request,
|
||||
wait_for_event,
|
||||
)
|
||||
|
||||
# ------------------------------
|
||||
# 配置常量(可根据需求调整)
|
||||
# ------------------------------
|
||||
DEFAULT_APPID = "7069844318"
|
||||
DEFAULT_ACCESS_TOKEN = "osFMEJr20SSTWRql43cJlZkAOg7iwvxu"
|
||||
DEFAULT_ENDPOINT = "wss://openspeech.bytedance.com/api/v3/tts/bidirection"
|
||||
DEFAULT_VOICE_TYPE = "zh_female_gaolengyujie_emo_v2_mars_bigtts"
|
||||
DEFAULT_ENCODING = "pcm"
|
||||
DEFAULT_SAMPLE_RATE = 16000
|
||||
|
||||
|
||||
@dataclass
|
||||
class TTSRequest:
|
||||
"""TTS请求对象(带唯一标识)"""
|
||||
tts_text: str
|
||||
mode: str = "预训练音色"
|
||||
sft_spk: str = ""
|
||||
seed: int = field(default_factory=lambda: np.random.randint(1, 100000000))
|
||||
stream: bool = True
|
||||
speed: float = 1.0
|
||||
prompt_wav: str = ""
|
||||
instruct_text: str = ""
|
||||
voice_type: str = DEFAULT_VOICE_TYPE
|
||||
encoding: str = DEFAULT_ENCODING
|
||||
speed: float = 1.0 # 语速(字节跳动TTS支持,需服务端兼容)
|
||||
stream: bool = True # 是否流式合成
|
||||
request_id: str = field(default_factory=lambda: str(uuid.uuid4())) # 唯一请求ID
|
||||
session_id: str = field(default_factory=lambda: str(uuid.uuid4())) # 会话ID(每个请求一个会话)
|
||||
|
||||
|
||||
class CosyVoiceTTSSocketClient:
|
||||
"""CosyVoice TTS WebSocket 客户端(异步/流式/带任务队列)"""
|
||||
class ByteDanceTTSSocketClient:
|
||||
"""字节跳动 TTS WebSocket 客户端(异步/流式/带任务队列)"""
|
||||
|
||||
def __init__(self, ws_url: str = "ws://localhost:50000/ws/tts", max_queue_size: int = 100):
|
||||
def __init__(
|
||||
self,
|
||||
appid: str = DEFAULT_APPID,
|
||||
access_token: str = DEFAULT_ACCESS_TOKEN,
|
||||
endpoint: str = DEFAULT_ENDPOINT,
|
||||
max_queue_size: int = 100
|
||||
):
|
||||
"""
|
||||
初始化客户端
|
||||
:param ws_url: WebSocket 服务端地址
|
||||
:param appid: 字节跳动APP ID
|
||||
:param access_token: 访问令牌
|
||||
:param endpoint: WebSocket 服务端地址
|
||||
:param max_queue_size: 最大队列长度(防止内存溢出)
|
||||
"""
|
||||
self.ws_url = ws_url
|
||||
# 基础配置
|
||||
self.appid = appid
|
||||
self.access_token = access_token
|
||||
self.endpoint = endpoint
|
||||
self.max_queue_size = max_queue_size
|
||||
|
||||
# WebSocket 连接状态
|
||||
self.websocket: Optional[websockets.WebSocketClientProtocol] = None
|
||||
self.is_connected = False
|
||||
self.is_processing = False # 是否正在处理请求
|
||||
self.logid: Optional[str] = None # 服务端返回的日志ID
|
||||
|
||||
# 异步任务队列(FIFO)
|
||||
self.request_queue: asyncio.Queue[TTSRequest] = asyncio.Queue(maxsize=max_queue_size)
|
||||
@@ -42,45 +77,84 @@ class CosyVoiceTTSSocketClient:
|
||||
# 回调函数定义(所有回调都带request_id,方便关联请求)
|
||||
self.on_task_enqueue: Callable[[str], None] = lambda req_id: None # 任务入队回调
|
||||
self.on_start: Callable[[str, Dict[str, Any]], None] = lambda req_id, data: None # 合成开始回调
|
||||
self.on_audio_chunk: Callable[[str, np.ndarray], None] = lambda req_id, chunk: None # 音频块回调
|
||||
self.on_audio_chunk: Callable[[str, bytes], None] = lambda req_id, chunk: None # 音频块回调(原始字节)
|
||||
self.on_end: Callable[[str, Dict[str, Any]], None] = lambda req_id, data: None # 合成结束回调
|
||||
self.on_error: Callable[[str, str], None] = lambda req_id, msg: None # 错误回调
|
||||
self.on_queue_full: Callable[[str], None] = lambda req_id: None # 队列满回调
|
||||
|
||||
# 音频播放流
|
||||
# 音频播放相关(支持MP3格式直接播放)
|
||||
self.play_stream: Optional[sd.OutputStream] = None
|
||||
self.current_req_id: Optional[str] = None
|
||||
self.enable_playback: bool = True # 是否启用实时播放
|
||||
|
||||
# 是否启用播放功能
|
||||
self.enable_playback: bool = True
|
||||
|
||||
# 外部回调函数,用于返回完整结果
|
||||
# 外部回调函数(返回完整结果)
|
||||
self.external_callback: Optional[Callable[[str, Dict[str, Any]], None]] = None
|
||||
|
||||
# 存储每个请求的音频数据
|
||||
self.audio_buffers: Dict[str, List[np.ndarray]] = {}
|
||||
# 存储每个请求的完整音频数据(原始字节)
|
||||
self.audio_buffers: Dict[str, List[bytes]] = {}
|
||||
|
||||
def _get_resource_id(self, voice_type: str) -> str:
|
||||
"""根据音色类型获取资源ID(字节跳动TTS协议要求)"""
|
||||
if voice_type.startswith("S_"):
|
||||
return "volc.megatts.default"
|
||||
return "volc.service_type.10029"
|
||||
|
||||
async def _create_websocket_connection(self):
|
||||
"""创建WebSocket连接(内部使用)"""
|
||||
headers = {
|
||||
"X-Api-App-Key": self.appid,
|
||||
"X-Api-Access-Key": self.access_token,
|
||||
"X-Api-Resource-Id": self._get_resource_id(DEFAULT_VOICE_TYPE), # 用默认音色获取资源ID
|
||||
"X-Api-Connect-Id": str(uuid.uuid4()),
|
||||
}
|
||||
|
||||
print(f"连接到 TTS 服务端: {self.endpoint}")
|
||||
self.websocket = await websockets.connect(
|
||||
self.endpoint,
|
||||
additional_headers=headers,
|
||||
max_size=10 * 1024 * 1024 # 10MB缓冲区
|
||||
)
|
||||
self.is_connected = True
|
||||
self.logid = self.websocket.response.headers.get("x-tt-logid")
|
||||
print(f"连接成功,LogID: {self.logid}")
|
||||
|
||||
# 发送连接启动指令
|
||||
await start_connection(self.websocket)
|
||||
await wait_for_event(
|
||||
self.websocket, MsgType.FullServerResponse, EventType.ConnectionStarted
|
||||
)
|
||||
print("TTS连接已初始化完成")
|
||||
|
||||
async def connect(self):
|
||||
"""建立 WebSocket 连接(初始化一次)"""
|
||||
"""建立WebSocket连接(外部调用,初始化一次)"""
|
||||
if not self.is_connected:
|
||||
try:
|
||||
self.websocket = await websockets.connect(self.ws_url)
|
||||
self.is_connected = True
|
||||
print(f"成功连接到 TTS 服务端: {self.ws_url}")
|
||||
|
||||
await self._create_websocket_connection()
|
||||
# 启动队列消费协程(后台运行)
|
||||
asyncio.create_task(self._consume_queue())
|
||||
except Exception as e:
|
||||
raise ConnectionError(f"连接失败: {str(e)}")
|
||||
error_msg = f"连接失败: {str(e)}"
|
||||
print(error_msg)
|
||||
raise ConnectionError(error_msg)
|
||||
|
||||
async def disconnect(self):
|
||||
"""关闭 WebSocket 连接"""
|
||||
"""关闭WebSocket连接"""
|
||||
if self.is_connected and self.websocket:
|
||||
await self.websocket.close()
|
||||
self.is_connected = False
|
||||
self.websocket = None
|
||||
print("已断开与 TTS 服务端的连接")
|
||||
try:
|
||||
# 发送连接结束指令
|
||||
await finish_connection(self.websocket)
|
||||
await wait_for_event(
|
||||
self.websocket, MsgType.FullServerResponse, EventType.ConnectionFinished
|
||||
)
|
||||
except Exception as e:
|
||||
print(f"关闭连接时异常: {str(e)}")
|
||||
finally:
|
||||
await self.websocket.close()
|
||||
self.is_connected = False
|
||||
self.websocket = None
|
||||
print("已断开与TTS服务端的连接")
|
||||
|
||||
# 清理音频播放流
|
||||
if self.play_stream:
|
||||
self.play_stream.stop()
|
||||
self.play_stream.close()
|
||||
@@ -91,281 +165,377 @@ class CosyVoiceTTSSocketClient:
|
||||
self.external_callback = callback
|
||||
|
||||
def set_playback_enabled(self, enabled: bool):
|
||||
"""设置是否启用音频播放"""
|
||||
"""设置是否启用音频实时播放"""
|
||||
self.enable_playback = enabled
|
||||
print(f"音频实时播放已{'启用' if enabled else '禁用'}")
|
||||
|
||||
async def add_tts_request(self, tts_text: str, **kwargs) -> str:
|
||||
async def synthesize(self, tts_text: str, **kwargs) -> str:
|
||||
"""
|
||||
添加TTS请求到队列(异步非阻塞)
|
||||
异步非阻塞添加TTS请求到队列
|
||||
:param tts_text: 要合成的文本
|
||||
:param kwargs: 其他TTS参数
|
||||
:param kwargs: 其他TTS参数(voice_type, encoding, speed等)
|
||||
:return: 唯一请求ID
|
||||
"""
|
||||
# 创建请求对象
|
||||
# 创建请求对象(支持覆盖默认参数)
|
||||
request = TTSRequest(tts_text=tts_text, **kwargs)
|
||||
req_id = request.request_id
|
||||
|
||||
# 初始化音频缓冲区
|
||||
self.audio_buffers[req_id] = []
|
||||
|
||||
# 异步入队
|
||||
# 异步入队(非阻塞)
|
||||
try:
|
||||
await self.request_queue.put(request)
|
||||
self.on_task_enqueue(req_id)
|
||||
print(f"请求 {req_id} 已加入队列,当前队列长度: {self.request_queue.qsize()}")
|
||||
print(f"请求 [{req_id[:8]}] 已加入队列,当前队列长度: {self.request_queue.qsize()}")
|
||||
return req_id
|
||||
except asyncio.QueueFull:
|
||||
self.on_queue_full(req_id)
|
||||
raise Exception(f"队列已满,请求 {req_id} 入队失败")
|
||||
error_msg = f"队列已满(最大长度{self.max_queue_size}),请求 [{req_id[:8]}] 入队失败"
|
||||
print(error_msg)
|
||||
raise Exception(error_msg)
|
||||
|
||||
async def _consume_queue(self):
|
||||
"""消费队列(后台协程,自动处理排队请求)"""
|
||||
print("队列消费协程已启动")
|
||||
while True:
|
||||
try:
|
||||
# 等待队列中有请求(阻塞)
|
||||
# 等待队列中有请求(阻塞,直到有任务)
|
||||
request = await self.request_queue.get()
|
||||
req_id = request.request_id
|
||||
|
||||
# 标记为处理中
|
||||
self.is_processing = True
|
||||
print(f"开始处理请求 {req_id},剩余队列长度: {self.request_queue.qsize()}")
|
||||
# print(f"\n开始处理请求 [{req_id[:8]}],剩余队列长度: {self.request_queue.qsize()}")
|
||||
|
||||
# 处理当前请求
|
||||
# 处理单个请求
|
||||
await self._process_single_request(request)
|
||||
|
||||
# 标记任务完成
|
||||
# 标记任务完成(让Queue知道可以继续)
|
||||
self.request_queue.task_done()
|
||||
self.is_processing = False
|
||||
|
||||
except Exception as e:
|
||||
print(f"队列消费异常: {str(e)}")
|
||||
error_msg = f"队列消费异常: {str(e)}"
|
||||
print(error_msg)
|
||||
self.is_processing = False
|
||||
# 短暂等待后继续消费,避免死循环
|
||||
# 短暂等待,避免死循环占用CPU
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
async def _process_single_request(self, request: TTSRequest):
|
||||
"""处理单个TTS请求"""
|
||||
req_id = request.request_id
|
||||
|
||||
# 参数校验
|
||||
if not request.tts_text:
|
||||
self.on_error(req_id, "合成文本不能为空")
|
||||
return
|
||||
|
||||
# 确保已连接
|
||||
if not self.is_connected:
|
||||
try:
|
||||
await self.connect()
|
||||
except Exception as e:
|
||||
self.on_error(req_id, f"连接失败: {str(e)}")
|
||||
return
|
||||
|
||||
# 构造请求数据
|
||||
request_data = {
|
||||
"tts_text": request.tts_text,
|
||||
"mode": request.mode,
|
||||
"sft_spk": request.sft_spk,
|
||||
"seed": request.seed,
|
||||
"stream": request.stream,
|
||||
"speed": request.speed,
|
||||
"prompt_wav": request.prompt_wav,
|
||||
"instruct_text": request.instruct_text
|
||||
def _build_base_request(self, request: TTSRequest) -> Dict[str, Any]:
|
||||
"""构建字节跳动TTS基础请求参数"""
|
||||
aaa = {
|
||||
"user": {"uid": str(uuid.uuid4())},
|
||||
"namespace": "BidirectionalTTS",
|
||||
"req_params": {
|
||||
"speaker": request.voice_type,
|
||||
"audio_params": {
|
||||
"format": request.encoding,
|
||||
"sample_rate": DEFAULT_SAMPLE_RATE,
|
||||
"enable_timestamp": True,
|
||||
},
|
||||
"additions": json.dumps({"disable_markdown_filter": False}),
|
||||
"speed": request.speed, # 语速参数(需服务端支持)
|
||||
},
|
||||
}
|
||||
print('aaa', aaa)
|
||||
return aaa
|
||||
async def _send_text_stream(self, request: TTSRequest, session_id: str):
|
||||
"""流式发送文本(逐字符发送,字节跳动TTS流式协议要求)"""
|
||||
base_request = self._build_base_request(request)
|
||||
text = request.tts_text.strip()
|
||||
|
||||
try:
|
||||
# 发送请求
|
||||
await self.websocket.send(json.dumps(request_data))
|
||||
|
||||
# 处理返回的流式数据
|
||||
await self._handle_response(req_id)
|
||||
|
||||
except websockets.exceptions.ConnectionClosed:
|
||||
self.is_connected = False
|
||||
self.on_error(req_id, "连接已关闭")
|
||||
except Exception as e:
|
||||
self.on_error(req_id, f"处理请求失败: {str(e)}")
|
||||
|
||||
async def _handle_response(self, req_id: str):
|
||||
"""处理单个请求的服务端响应"""
|
||||
if not self.websocket:
|
||||
if not text:
|
||||
print(f"请求 [{request.request_id[:8]}] 文本为空,跳过发送")
|
||||
return
|
||||
|
||||
# 逐字符发送(控制发送速率,避免拥塞)
|
||||
for char in text:
|
||||
if not self.is_connected or not self.websocket:
|
||||
raise ConnectionError("连接已断开,无法继续发送文本")
|
||||
|
||||
# 构建单个字符的任务请求
|
||||
task_req = copy.deepcopy(base_request)
|
||||
task_req["event"] = EventType.TaskRequest
|
||||
task_req["req_params"]["text"] = char
|
||||
|
||||
# 发送任务请求
|
||||
await task_request(
|
||||
self.websocket,
|
||||
json.dumps(task_req).encode("utf-8"),
|
||||
session_id
|
||||
)
|
||||
# 控制发送速率(5ms/字符,可调整)
|
||||
await asyncio.sleep(0.005)
|
||||
|
||||
# 发送会话结束指令
|
||||
await finish_session(self.websocket, session_id)
|
||||
print(f"请求 [{request.request_id[:8]}] 文本发送完成")
|
||||
|
||||
async def _handle_audio_response(self, req_id: str, session_id: str, request: TTSRequest) -> Dict[str, Any]:
|
||||
"""处理服务端的流式音频响应"""
|
||||
if not self.websocket:
|
||||
raise ConnectionError("WebSocket连接未建立")
|
||||
|
||||
sample_rate = DEFAULT_SAMPLE_RATE
|
||||
audio_received = False
|
||||
|
||||
try:
|
||||
sample_rate = None
|
||||
while True:
|
||||
# 接收服务端消息(异步)
|
||||
response = await self.websocket.recv()
|
||||
data = json.loads(response)
|
||||
# 接收服务端消息(异步阻塞)
|
||||
msg = await receive_message(self.websocket)
|
||||
|
||||
# 根据状态分发到不同回调(都带req_id)
|
||||
status = data.get("status")
|
||||
if status == "start":
|
||||
# 合成开始 - 返回采样率等信息
|
||||
sample_rate = data.get("sample_rate")
|
||||
self.on_start(req_id, data)
|
||||
if msg.type == MsgType.FullServerResponse:
|
||||
# 完整响应(开始/结束/错误)
|
||||
if msg.event == EventType.SessionStarted:
|
||||
# 会话开始回调
|
||||
start_data = {
|
||||
"session_id": session_id,
|
||||
"sample_rate": sample_rate,
|
||||
"encoding": request.encoding,
|
||||
"voice_type": request.voice_type,
|
||||
"logid": self.logid
|
||||
}
|
||||
self.on_start(req_id, start_data)
|
||||
print(f"请求 [{req_id[:8]}] 合成开始")
|
||||
|
||||
elif status == "stream":
|
||||
# 流式音频块 - 转为 numpy 数组
|
||||
audio_chunk = np.array(data["audio_chunk"], dtype=np.float32)
|
||||
elif msg.event == EventType.SessionFinished:
|
||||
# 会话结束,退出循环
|
||||
end_data = {"session_id": session_id, "message": "合成完成"}
|
||||
self.on_end(req_id, end_data)
|
||||
print(f"请求 [{req_id[:8]}] 合成结束")
|
||||
break
|
||||
|
||||
# 保存到缓冲区
|
||||
self.audio_buffers[req_id].append(audio_chunk)
|
||||
elif msg.type == MsgType.AudioOnlyServer:
|
||||
# 流式音频数据(原始字节)
|
||||
audio_chunk = msg.payload
|
||||
if audio_chunk:
|
||||
audio_received = True
|
||||
# 保存到缓冲区
|
||||
self.audio_buffers[req_id].append(audio_chunk)
|
||||
# 音频块回调
|
||||
self.on_audio_chunk(req_id, audio_chunk)
|
||||
# 实时播放(如果启用)
|
||||
await self._play_audio_chunk(req_id, audio_chunk, sample_rate)
|
||||
|
||||
# 音频块回调
|
||||
self.on_audio_chunk(req_id, audio_chunk)
|
||||
else:
|
||||
# 未知消息类型
|
||||
raise RuntimeError(f"收到未知消息类型: {msg.type}, 内容: {msg}")
|
||||
|
||||
elif status == "end":
|
||||
# 合成结束
|
||||
self.on_end(req_id, data)
|
||||
|
||||
# 组装完整音频数据
|
||||
full_audio = np.concatenate(self.audio_buffers[req_id]) if self.audio_buffers[req_id] else np.array(
|
||||
[], dtype=np.float32)
|
||||
|
||||
# 准备返回给外部回调的数据
|
||||
result_data = {
|
||||
"status": "completed",
|
||||
"request_id": req_id,
|
||||
"sample_rate": sample_rate,
|
||||
"audio_data": full_audio,
|
||||
"audio_length": len(full_audio),
|
||||
"message": data.get("msg", "合成完成")
|
||||
}
|
||||
|
||||
# 调用外部回调(非阻塞)
|
||||
if self.external_callback:
|
||||
# 使用create_task确保回调不会阻塞主流程
|
||||
asyncio.create_task(self._call_external_callback(req_id, result_data))
|
||||
|
||||
# 清理缓冲区
|
||||
if req_id in self.audio_buffers:
|
||||
del self.audio_buffers[req_id]
|
||||
|
||||
break
|
||||
|
||||
elif status == "error":
|
||||
# 错误处理
|
||||
error_msg = data.get("msg", "未知错误")
|
||||
self.on_error(req_id, error_msg)
|
||||
|
||||
# 准备错误结果数据
|
||||
error_data = {
|
||||
"status": "error",
|
||||
"request_id": req_id,
|
||||
"message": error_msg
|
||||
}
|
||||
|
||||
# 调用外部回调
|
||||
if self.external_callback:
|
||||
asyncio.create_task(self._call_external_callback(req_id, error_data))
|
||||
|
||||
# 清理缓冲区
|
||||
if req_id in self.audio_buffers:
|
||||
del self.audio_buffers[req_id]
|
||||
|
||||
break
|
||||
# 组装完整结果
|
||||
full_audio = b"".join(self.audio_buffers[req_id]) if self.audio_buffers[req_id] else b""
|
||||
return {
|
||||
"status": "completed",
|
||||
"request_id": req_id,
|
||||
"session_id": session_id,
|
||||
"sample_rate": sample_rate,
|
||||
"encoding": request.encoding,
|
||||
"audio_data": full_audio, # 完整音频字节数据
|
||||
"audio_length": len(full_audio),
|
||||
"message": "合成成功" if audio_received else "合成完成但未收到音频数据"
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
error_msg = f"接收数据异常: {str(e)}"
|
||||
error_msg = f"处理音频响应异常: {str(e)}"
|
||||
self.on_error(req_id, error_msg)
|
||||
|
||||
# 准备错误结果数据
|
||||
error_data = {
|
||||
return {
|
||||
"status": "error",
|
||||
"request_id": req_id,
|
||||
"session_id": session_id,
|
||||
"message": error_msg
|
||||
}
|
||||
|
||||
# 调用外部回调
|
||||
if self.external_callback:
|
||||
asyncio.create_task(self._call_external_callback(req_id, error_data))
|
||||
async def _play_audio_chunk(self, req_id: str, chunk: bytes, sample_rate: int):
|
||||
"""实时播放音频块(支持MP3格式)"""
|
||||
if not self.enable_playback:
|
||||
return
|
||||
|
||||
# 确保当前请求是正在播放的请求
|
||||
if self.current_req_id is None:
|
||||
self.current_req_id = req_id
|
||||
|
||||
if req_id != self.current_req_id:
|
||||
# 切换请求时,重置播放流
|
||||
if self.play_stream:
|
||||
self.play_stream.stop()
|
||||
self.play_stream.close()
|
||||
self.current_req_id = req_id
|
||||
|
||||
try:
|
||||
# 初始化播放流(如果未初始化)
|
||||
if not self.play_stream:
|
||||
self.play_stream = sd.OutputStream(
|
||||
samplerate=sample_rate,
|
||||
channels=1, # 单声道
|
||||
dtype=np.float32
|
||||
)
|
||||
self.play_stream.start()
|
||||
|
||||
# MP3字节 → 音频数组(直接播放)
|
||||
# 注意:sounddevice默认支持PCM格式,如果是MP3需要解码,这里简化处理(实际使用建议用pydub解码)
|
||||
# 如需支持MP3播放,请安装 pydub: pip install pydub ffmpeg
|
||||
try:
|
||||
# 简化处理:假设服务端返回PCM(如果是MP3,需替换为解码逻辑)
|
||||
audio_array = np.frombuffer(chunk, dtype=np.float32)
|
||||
if audio_array.size > 0:
|
||||
self.play_stream.write(audio_array)
|
||||
except Exception as e:
|
||||
print(f"音频播放异常: {str(e)},请确保音频格式正确")
|
||||
|
||||
except Exception as e:
|
||||
print(f"播放流初始化失败: {str(e)}")
|
||||
|
||||
async def _process_single_request(self, request: TTSRequest):
|
||||
"""处理单个TTS请求(完整流程:连接→启动会话→流式发送文本→接收音频→回调结果)"""
|
||||
req_id = request.request_id
|
||||
session_id = request.session_id
|
||||
|
||||
# 参数校验
|
||||
if not request.tts_text.strip():
|
||||
error_msg = "合成文本不能为空"
|
||||
self.on_error(req_id, error_msg)
|
||||
self._send_external_callback(req_id, {
|
||||
"status": "error",
|
||||
"request_id": req_id,
|
||||
"session_id": session_id,
|
||||
"message": error_msg
|
||||
})
|
||||
return
|
||||
|
||||
# 确保连接已建立(断开时自动重连)
|
||||
if not self.is_connected:
|
||||
print(f"请求 [{req_id[:8]}] 处理时连接已断开,尝试重连...")
|
||||
try:
|
||||
await self._create_websocket_connection()
|
||||
except Exception as e:
|
||||
error_msg = f"重连失败: {str(e)}"
|
||||
self.on_error(req_id, error_msg)
|
||||
self._send_external_callback(req_id, {
|
||||
"status": "error",
|
||||
"request_id": req_id,
|
||||
"session_id": session_id,
|
||||
"message": error_msg
|
||||
})
|
||||
return
|
||||
|
||||
try:
|
||||
# 1. 启动会话
|
||||
base_request = self._build_base_request(request)
|
||||
start_session_req = copy.deepcopy(base_request)
|
||||
start_session_req["event"] = EventType.StartSession
|
||||
|
||||
await start_session(
|
||||
self.websocket,
|
||||
json.dumps(start_session_req).encode("utf-8"),
|
||||
session_id
|
||||
)
|
||||
# 等待会话启动成功
|
||||
await wait_for_event(
|
||||
self.websocket, MsgType.FullServerResponse, EventType.SessionStarted
|
||||
)
|
||||
|
||||
# 2. 异步流式发送文本(后台任务,不阻塞接收音频)
|
||||
send_task = asyncio.create_task(self._send_text_stream(request, session_id))
|
||||
|
||||
# 3. 接收并处理音频响应
|
||||
result_data = await self._handle_audio_response(req_id, session_id, request)
|
||||
|
||||
# 4. 等待文本发送任务完成
|
||||
await send_task
|
||||
|
||||
# 5. 发送外部回调
|
||||
self._send_external_callback(req_id, result_data)
|
||||
|
||||
except Exception as e:
|
||||
error_msg = f"处理请求 [{req_id[:8]}] 异常: {str(e)}"
|
||||
self.on_error(req_id, error_msg)
|
||||
self._send_external_callback(req_id, {
|
||||
"status": "error",
|
||||
"request_id": req_id,
|
||||
"session_id": session_id,
|
||||
"message": error_msg
|
||||
})
|
||||
finally:
|
||||
# 清理缓冲区
|
||||
if req_id in self.audio_buffers:
|
||||
del self.audio_buffers[req_id]
|
||||
|
||||
async def _call_external_callback(self, req_id: str, result_data: Dict[str, Any]):
|
||||
"""异步调用外部回调函数"""
|
||||
try:
|
||||
if self.external_callback:
|
||||
# 确保回调是异步的,不会阻塞
|
||||
if asyncio.iscoroutinefunction(self.external_callback):
|
||||
await self.external_callback(req_id, result_data)
|
||||
else:
|
||||
# 如果是同步函数,在线程池中执行
|
||||
await asyncio.get_event_loop().run_in_executor(
|
||||
None, self.external_callback, req_id, result_data
|
||||
)
|
||||
except Exception as e:
|
||||
print(f"外部回调执行异常: {str(e)}")
|
||||
|
||||
def _play_audio_chunk(self, chunk: np.ndarray):
|
||||
"""播放音频块(非阻塞)"""
|
||||
if not self.enable_playback or not self.play_stream:
|
||||
def _send_external_callback(self, req_id: str, result_data: Dict[str, Any]):
|
||||
"""发送外部回调(支持同步/异步回调函数)"""
|
||||
if not self.external_callback:
|
||||
return
|
||||
|
||||
try:
|
||||
if chunk.size > 0:
|
||||
self.play_stream.write(chunk)
|
||||
# 异步回调:直接await
|
||||
if asyncio.iscoroutinefunction(self.external_callback):
|
||||
asyncio.create_task(self.external_callback(req_id, result_data))
|
||||
# 同步回调:在线程池中执行(避免阻塞事件循环)
|
||||
else:
|
||||
asyncio.get_event_loop().run_in_executor(
|
||||
None, self.external_callback, req_id, result_data
|
||||
)
|
||||
except Exception as e:
|
||||
print(f"音频播放异常: {str(e)}")
|
||||
print(f"外部回调执行异常: {str(e)}")
|
||||
|
||||
async def wait_queue_empty(self):
|
||||
"""等待队列所有任务处理完成(阻塞)"""
|
||||
async def wait_all_completed(self):
|
||||
"""等待队列中所有任务处理完成(阻塞)"""
|
||||
await self.request_queue.join()
|
||||
print("所有队列任务已处理完成")
|
||||
print("\n所有队列任务已处理完成")
|
||||
|
||||
|
||||
# ------------------------------
|
||||
# 使用示例
|
||||
# 使用示例(与你提供的风格完全一致)
|
||||
# ------------------------------
|
||||
class TTSManager:
|
||||
"""TTS管理器 - 供外部代码调用"""
|
||||
"""TTS管理器 - 供外部代码调用(封装客户端,简化使用)"""
|
||||
|
||||
def __init__(self, ws_url: str = "ws://localhost:50000/ws/tts"):
|
||||
self.client = CosyVoiceTTSSocketClient(ws_url)
|
||||
self._setup_callbacks()
|
||||
def __init__(
|
||||
self,
|
||||
appid: str = DEFAULT_APPID,
|
||||
access_token: str = DEFAULT_ACCESS_TOKEN,
|
||||
endpoint: str = DEFAULT_ENDPOINT
|
||||
):
|
||||
self.client = ByteDanceTTSSocketClient(
|
||||
appid=appid,
|
||||
access_token=access_token,
|
||||
endpoint=endpoint
|
||||
)
|
||||
self._setup_internal_callbacks()
|
||||
|
||||
def _setup_callbacks(self):
|
||||
"""设置内部回调"""
|
||||
def _setup_internal_callbacks(self):
|
||||
"""设置内部回调(日志/状态提示)"""
|
||||
|
||||
# 合成开始回调
|
||||
def on_tts_start(req_id, data):
|
||||
print(f"\n🎤 开始合成 [{req_id[:8]}] - 采样率: {data.get('sample_rate')}")
|
||||
def on_task_enqueue(req_id: str):
|
||||
"""任务入队回调"""
|
||||
print(f"📥 任务 [{req_id[:8]}] 已入队")
|
||||
|
||||
# 初始化播放流(如果需要播放)
|
||||
if self.client.enable_playback:
|
||||
sample_rate = data.get("sample_rate", 24000)
|
||||
if self.client.play_stream:
|
||||
self.client.play_stream.stop()
|
||||
self.client.play_stream.close()
|
||||
def on_tts_start(req_id: str, data: Dict[str, Any]):
|
||||
"""合成开始回调"""
|
||||
print(f"🎤 合成开始 [{req_id[:8]}] - 采样率: {data['sample_rate']}, 编码: {data['encoding']}")
|
||||
|
||||
self.client.play_stream = sd.OutputStream(
|
||||
samplerate=sample_rate,
|
||||
channels=1,
|
||||
dtype=np.float32
|
||||
)
|
||||
self.client.play_stream.start()
|
||||
self.client.current_req_id = req_id
|
||||
def on_audio_chunk(req_id: str, chunk: bytes):
|
||||
"""音频块回调(内部仅打印日志,外部通过external_callback获取)"""
|
||||
print(f"🔊 收到音频块 [{req_id[:8]}] - 大小: {len(chunk)}字节", end="\r")
|
||||
|
||||
# 音频块回调
|
||||
def on_audio_chunk(req_id, chunk):
|
||||
if req_id == self.client.current_req_id:
|
||||
self.client._play_audio_chunk(chunk)
|
||||
def on_tts_end(req_id: str, data: Dict[str, Any]):
|
||||
"""合成结束回调"""
|
||||
print(f"\n🏁 合成结束 [{req_id[:8]}] - 会话ID: {data['session_id']}")
|
||||
|
||||
# 合成结束回调
|
||||
def on_tts_end(req_id, data):
|
||||
print(f"🏁 合成完成 [{req_id[:8]}] - {data.get('msg', '完成')}")
|
||||
def on_tts_error(req_id: str, msg: str):
|
||||
"""错误回调"""
|
||||
print(f"\n❌ 合成失败 [{req_id[:8]}] - 错误: {msg}")
|
||||
|
||||
# 错误回调
|
||||
def on_tts_error(req_id, msg):
|
||||
print(f"❌ 合成失败 [{req_id[:8]}] - {msg}")
|
||||
def on_queue_full(req_id: str):
|
||||
"""队列满回调"""
|
||||
print(f"⚠️ 队列已满,请求 [{req_id[:8]}] 入队失败")
|
||||
|
||||
# 绑定内部回调
|
||||
self.client.on_task_enqueue = on_task_enqueue
|
||||
self.client.on_start = on_tts_start
|
||||
self.client.on_audio_chunk = on_audio_chunk
|
||||
self.client.on_end = on_tts_end
|
||||
self.client.on_error = on_tts_error
|
||||
self.client.on_queue_full = on_queue_full
|
||||
|
||||
async def initialize(self):
|
||||
"""初始化连接"""
|
||||
@@ -376,25 +546,25 @@ class TTSManager:
|
||||
await self.client.disconnect()
|
||||
|
||||
def set_result_callback(self, callback: Callable[[str, Dict[str, Any]], None]):
|
||||
"""设置结果回调函数"""
|
||||
"""设置外部结果回调(获取完整音频数据)"""
|
||||
self.client.set_external_callback(callback)
|
||||
|
||||
def set_playback_enabled(self, enabled: bool):
|
||||
"""设置是否启用播放"""
|
||||
"""设置是否启用实时播放"""
|
||||
self.client.set_playback_enabled(enabled)
|
||||
|
||||
async def synthesize(self, text: str, **kwargs) -> str:
|
||||
"""
|
||||
异步合成文本(非阻塞)
|
||||
异步非阻塞合成文本
|
||||
:param text: 要合成的文本
|
||||
:param kwargs: 其他参数(mode, sft_spk, speed等)
|
||||
:param kwargs: 其他参数(voice_type, encoding, speed等)
|
||||
:return: 请求ID
|
||||
"""
|
||||
return await self.client.add_tts_request(text, **kwargs)
|
||||
return await self.client.synthesize(text, **kwargs)
|
||||
|
||||
async def wait_all_completed(self):
|
||||
"""等待所有任务完成"""
|
||||
await self.client.wait_queue_empty()
|
||||
await self.client.wait_all_completed()
|
||||
|
||||
|
||||
# ------------------------------
|
||||
@@ -402,147 +572,120 @@ class TTSManager:
|
||||
# ------------------------------
|
||||
async def external_usage_example():
|
||||
"""外部代码使用示例"""
|
||||
# 1. 创建TTS管理器(可替换为自己的appid和access_token)
|
||||
tts_manager = TTSManager(
|
||||
appid=DEFAULT_APPID,
|
||||
access_token=DEFAULT_ACCESS_TOKEN,
|
||||
endpoint=DEFAULT_ENDPOINT
|
||||
)
|
||||
|
||||
# 1. 创建TTS管理器
|
||||
tts_manager = TTSManager(ws_url="ws://10.10.10.202:50000/ws/tts")
|
||||
|
||||
# 2. 设置结果回调函数(接收完整结果)
|
||||
def handle_tts_result1(req_id: str, result: Dict[str, Any]):
|
||||
"""处理TTS结果回调"""
|
||||
# 2. 设置外部结果回调(获取完整音频数据)
|
||||
def handle_tts_result(req_id: str, result: Dict[str, Any]):
|
||||
"""处理TTS完整结果(同步回调)"""
|
||||
status = result.get("status")
|
||||
|
||||
if status == "completed":
|
||||
audio_data = result.get("audio_data")
|
||||
sample_rate = result.get("sample_rate")
|
||||
encoding = result.get("encoding")
|
||||
audio_length = result.get("audio_length")
|
||||
|
||||
print(f"✅ 收到TTS结果 [{req_id[:8]}]: 长度{audio_length}采样点, 采样率{sample_rate}Hz")
|
||||
print(f"\n✅ 收到完整结果 [{req_id[:8]}] - 长度: {audio_length}字节, 编码: {encoding}")
|
||||
|
||||
# 这里可以保存音频文件或进行其他处理
|
||||
# 注意:audio_data是完整的numpy数组
|
||||
# 保存音频文件
|
||||
filename = f"tts_output_{req_id[:8]}.{encoding}"
|
||||
with open(filename, "wb") as f:
|
||||
f.write(audio_data)
|
||||
print(f"💾 音频文件已保存: {filename}")
|
||||
|
||||
elif status == "error":
|
||||
error_msg = result.get("message")
|
||||
print(f"❌ TTS处理失败 [{req_id[:8]}]: {error_msg}")
|
||||
print(f"\n❌ 请求 [{req_id[:8]}] 处理失败: {error_msg}")
|
||||
|
||||
def handle_tts_result(req_id: str, result: Dict[str, Any]):
|
||||
"""处理TTS结果回调 - 保存为PCM文件"""
|
||||
print('处理TTS结果回调', result)
|
||||
status = result.get("status")
|
||||
print('status')
|
||||
if status == "completed":
|
||||
audio_data = result.get("audio_data")
|
||||
sample_rate = result.get("sample_rate")
|
||||
|
||||
if audio_data is not None and len(audio_data) > 0:
|
||||
# 转换为PCM数据
|
||||
pcm_data = (audio_data.astype(np.float32) * 32767).astype(np.int16)
|
||||
pcm_bytes = pcm_data.tobytes()
|
||||
|
||||
# 保存为PCM文件
|
||||
filename = f"tts_output_{req_id[:8]}.pcm"
|
||||
with open(filename, "wb") as f:
|
||||
f.write(pcm_bytes)
|
||||
|
||||
print(f"💾 PCM文件已保存: {filename}, 大小: {len(pcm_bytes)}字节")
|
||||
|
||||
return filename
|
||||
|
||||
elif status == "error":
|
||||
error_msg = result.get("message")
|
||||
print(f"❌ TTS处理失败 [{req_id[:8]}]: {error_msg}")
|
||||
|
||||
return None
|
||||
# 绑定外部回调
|
||||
tts_manager.set_result_callback(handle_tts_result)
|
||||
|
||||
# 3. 设置是否播放(可选,默认True)
|
||||
tts_manager.set_playback_enabled(True) # 设置为False则不播放
|
||||
# 3. 设置是否启用实时播放(默认True)
|
||||
tts_manager.set_playback_enabled(True)
|
||||
|
||||
# 4. 初始化连接
|
||||
await tts_manager.initialize()
|
||||
|
||||
# 5. 异步合成多个文本(非阻塞)
|
||||
# 5. 异步提交多个TTS请求(非阻塞)
|
||||
texts = [
|
||||
"你好,这是第一个排队的TTS请求。",
|
||||
"我是第二个请求,会等第一个处理完再执行。",
|
||||
"第三个请求,支持流式播放和队列管理。",
|
||||
"第四个请求,测试队列的自动消费功能。",
|
||||
"最后一个请求,处理完成后会自动结束。"
|
||||
"你好,这是字节跳动TTS的流式合成测试。",
|
||||
"我支持异步非阻塞调用,多个请求可以排队处理。",
|
||||
"每个请求都会返回唯一的ID,方便你跟踪结果。",
|
||||
"音频数据会通过回调函数返回,支持实时播放和保存文件。",
|
||||
"最后一个测试句子,演示队列的自动消费功能。"
|
||||
]
|
||||
|
||||
req_ids = []
|
||||
for i, text in enumerate(texts):
|
||||
# 提交请求(非阻塞,立即返回)
|
||||
req_id = await tts_manager.synthesize(
|
||||
text,
|
||||
mode="预训练音色",
|
||||
sft_spk="中文女",
|
||||
voice_type=DEFAULT_VOICE_TYPE,
|
||||
encoding=DEFAULT_ENCODING,
|
||||
speed=1.0
|
||||
)
|
||||
req_ids.append(req_id)
|
||||
print(f"已提交请求 {i + 1}: ID={req_id[:8]}")
|
||||
print(f"📤 已提交请求 {i+1}: ID={req_id[:8]}")
|
||||
|
||||
# 可以立即继续其他操作,不需要等待
|
||||
await asyncio.sleep(0.5) # 模拟其他操作
|
||||
# 模拟其他业务逻辑(无需等待TTS完成)
|
||||
await asyncio.sleep(0.3)
|
||||
|
||||
# 6. 可以在这里做其他事情,TTS会在后台处理
|
||||
|
||||
# 7. 等待所有TTS任务完成(可选)
|
||||
# 6. 等待所有TTS任务完成(可选,根据业务需求决定是否等待)
|
||||
await tts_manager.wait_all_completed()
|
||||
|
||||
# 8. 关闭连接
|
||||
# 7. 关闭连接(程序退出前调用)
|
||||
await tts_manager.shutdown()
|
||||
|
||||
|
||||
# ------------------------------
|
||||
# 异步结果回调示例(高级用法)
|
||||
# ------------------------------
|
||||
async def async_result_callback(req_id: str, result: Dict[str, Any]):
|
||||
"""异步结果回调(支持异步操作,如上传音频到服务器)"""
|
||||
if result["status"] == "completed":
|
||||
print(f"\n⚡ 异步处理结果 [{req_id[:8]}] - 开始上传音频...")
|
||||
# 模拟异步上传操作
|
||||
await asyncio.sleep(0.5)
|
||||
print(f"⚡ 异步处理结果 [{req_id[:8]}] - 音频上传完成")
|
||||
|
||||
|
||||
async def advanced_usage_example():
|
||||
"""高级使用示例:动态控制播放和结果处理"""
|
||||
"""高级使用示例:异步回调 + 禁用播放 + 批量请求"""
|
||||
tts_manager = TTSManager()
|
||||
|
||||
tts_manager = TTSManager(ws_url="ws://10.10.10.202:50000/ws/tts")
|
||||
# 设置异步结果回调
|
||||
tts_manager.set_result_callback(async_result_callback)
|
||||
|
||||
# 自定义结果处理器
|
||||
class ResultProcessor:
|
||||
def __init__(self):
|
||||
self.results = {}
|
||||
|
||||
async def process_result(self, req_id: str, result: Dict[str, Any]):
|
||||
"""异步处理结果"""
|
||||
if result.get("status") == "completed":
|
||||
# 保存音频数据等处理
|
||||
self.results[req_id] = result
|
||||
print(f"处理完成: {req_id[:8]}, 音频长度: {result.get('audio_length')}")
|
||||
|
||||
# 这里可以添加自定义逻辑,如保存到文件、发送到网络等
|
||||
|
||||
processor = ResultProcessor()
|
||||
tts_manager.set_result_callback(processor.process_result)
|
||||
|
||||
# 禁用播放,只获取数据
|
||||
# 禁用实时播放(只获取音频数据)
|
||||
tts_manager.set_playback_enabled(False)
|
||||
|
||||
await tts_manager.initialize()
|
||||
|
||||
# 批量合成
|
||||
# 批量提交请求(并行提交)
|
||||
tasks = []
|
||||
for i in range(3):
|
||||
text = f"这是第{i+1}个测试文本,用于演示高级功能。"
|
||||
task = tts_manager.synthesize(text)
|
||||
text = f"这是第{i+1}个高级测试文本,使用异步回调处理结果。"
|
||||
task = tts_manager.synthesize(text, speed=0.9)
|
||||
tasks.append(task)
|
||||
|
||||
# 并行提交所有请求
|
||||
req_ids = await asyncio.gather(*tasks)
|
||||
# 并行提交所有请求
|
||||
req_ids = await asyncio.gather(*tasks)
|
||||
print(f"\n已并行提交 {len(req_ids)} 个请求")
|
||||
|
||||
# 等待完成
|
||||
await tts_manager.wait_all_completed()
|
||||
await tts_manager.shutdown()
|
||||
# 等待所有任务完成
|
||||
await tts_manager.wait_all_completed()
|
||||
await tts_manager.shutdown()
|
||||
|
||||
print(f"所有任务完成,共处理 {len(processor.results)} 个结果")
|
||||
|
||||
if __name__ == "__main__":
|
||||
# 运行示例
|
||||
try:
|
||||
# 基础使用示例
|
||||
# 运行基础使用示例
|
||||
asyncio.run(external_usage_example())
|
||||
|
||||
# 高级使用示例(取消注释运行)
|
||||
# 运行高级使用示例(取消注释)
|
||||
# asyncio.run(advanced_usage_example())
|
||||
|
||||
except KeyboardInterrupt:
|
||||
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Reference in New Issue
Block a user