mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-29 16:53:57 +08:00
update:添加编码解码
This commit is contained in:
@@ -10,6 +10,7 @@ from core.handle.reportHandle import enqueue_tts_report
|
|||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from core.utils.tts import MarkdownCleaner
|
from core.utils.tts import MarkdownCleaner
|
||||||
from core.utils.util import audio_to_data
|
from core.utils.util import audio_to_data
|
||||||
|
from core.utils import opus_encoder_utils
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
@@ -23,6 +24,9 @@ class TTSProviderBase(ABC):
|
|||||||
self.output_file = config.get("output_dir")
|
self.output_file = config.get("output_dir")
|
||||||
self.tts_queue = queue.Queue()
|
self.tts_queue = queue.Queue()
|
||||||
self.audio_play_queue = queue.Queue()
|
self.audio_play_queue = queue.Queue()
|
||||||
|
self.opus_encoder = opus_encoder_utils.OpusEncoderUtils(
|
||||||
|
sample_rate=16000, channels=1, frame_size_ms=60
|
||||||
|
)
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def generate_filename(self):
|
def generate_filename(self):
|
||||||
@@ -177,3 +181,7 @@ class TTSProviderBase(ABC):
|
|||||||
logger.bind(tag=TAG).error(
|
logger.bind(tag=TAG).error(
|
||||||
f"audio_play_priority priority_thread: {text} {e}"
|
f"audio_play_priority priority_thread: {text} {e}"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def wav_to_opus_data_audio_raw(self, raw_data_var, is_end=False):
|
||||||
|
opus_datas = self.opus_encoder.encode_pcm_to_opus(raw_data_var, is_end)
|
||||||
|
return opus_datas
|
||||||
|
|||||||
@@ -0,0 +1,420 @@
|
|||||||
|
import asyncio
|
||||||
|
import io
|
||||||
|
import os
|
||||||
|
import threading
|
||||||
|
import traceback
|
||||||
|
import uuid
|
||||||
|
import json
|
||||||
|
from datetime import datetime
|
||||||
|
|
||||||
|
import websockets
|
||||||
|
|
||||||
|
from config.logger import setup_logging
|
||||||
|
from core.providers.tts.base import TTSProviderBase
|
||||||
|
|
||||||
|
TAG = __name__
|
||||||
|
logger = setup_logging()
|
||||||
|
|
||||||
|
PROTOCOL_VERSION = 0b0001
|
||||||
|
DEFAULT_HEADER_SIZE = 0b0001
|
||||||
|
|
||||||
|
# Message Type:
|
||||||
|
FULL_CLIENT_REQUEST = 0b0001
|
||||||
|
AUDIO_ONLY_RESPONSE = 0b1011
|
||||||
|
FULL_SERVER_RESPONSE = 0b1001
|
||||||
|
ERROR_INFORMATION = 0b1111
|
||||||
|
|
||||||
|
# Message Type Specific Flags
|
||||||
|
MsgTypeFlagNoSeq = 0b0000 # Non-terminal packet with no sequence
|
||||||
|
MsgTypeFlagPositiveSeq = 0b1 # Non-terminal packet with sequence > 0
|
||||||
|
MsgTypeFlagLastNoSeq = 0b10 # last packet with no sequence
|
||||||
|
MsgTypeFlagNegativeSeq = 0b11 # Payload contains event number (int32)
|
||||||
|
MsgTypeFlagWithEvent = 0b100
|
||||||
|
# Message Serialization
|
||||||
|
NO_SERIALIZATION = 0b0000
|
||||||
|
JSON = 0b0001
|
||||||
|
# Message Compression
|
||||||
|
COMPRESSION_NO = 0b0000
|
||||||
|
COMPRESSION_GZIP = 0b0001
|
||||||
|
|
||||||
|
EVENT_NONE = 0
|
||||||
|
EVENT_Start_Connection = 1
|
||||||
|
|
||||||
|
EVENT_FinishConnection = 2
|
||||||
|
|
||||||
|
EVENT_ConnectionStarted = 50 # 成功建连
|
||||||
|
|
||||||
|
EVENT_ConnectionFailed = 51 # 建连失败(可能是无法通过权限认证)
|
||||||
|
|
||||||
|
EVENT_ConnectionFinished = 52 # 连接结束
|
||||||
|
|
||||||
|
# 上行Session事件
|
||||||
|
EVENT_StartSession = 100
|
||||||
|
|
||||||
|
EVENT_FinishSession = 102
|
||||||
|
# 下行Session事件
|
||||||
|
EVENT_SessionStarted = 150
|
||||||
|
EVENT_SessionFinished = 152
|
||||||
|
|
||||||
|
EVENT_SessionFailed = 153
|
||||||
|
|
||||||
|
# 上行通用事件
|
||||||
|
EVENT_TaskRequest = 200
|
||||||
|
|
||||||
|
# 下行TTS事件
|
||||||
|
EVENT_TTSSentenceStart = 350
|
||||||
|
|
||||||
|
EVENT_TTSSentenceEnd = 351
|
||||||
|
|
||||||
|
EVENT_TTSResponse = 352
|
||||||
|
|
||||||
|
|
||||||
|
class Header:
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
protocol_version=PROTOCOL_VERSION,
|
||||||
|
header_size=DEFAULT_HEADER_SIZE,
|
||||||
|
message_type: int = 0,
|
||||||
|
message_type_specific_flags: int = 0,
|
||||||
|
serial_method: int = NO_SERIALIZATION,
|
||||||
|
compression_type: int = COMPRESSION_NO,
|
||||||
|
reserved_data=0,
|
||||||
|
):
|
||||||
|
self.header_size = header_size
|
||||||
|
self.protocol_version = protocol_version
|
||||||
|
self.message_type = message_type
|
||||||
|
self.message_type_specific_flags = message_type_specific_flags
|
||||||
|
self.serial_method = serial_method
|
||||||
|
self.compression_type = compression_type
|
||||||
|
self.reserved_data = reserved_data
|
||||||
|
|
||||||
|
def as_bytes(self) -> bytes:
|
||||||
|
return bytes(
|
||||||
|
[
|
||||||
|
(self.protocol_version << 4) | self.header_size,
|
||||||
|
(self.message_type << 4) | self.message_type_specific_flags,
|
||||||
|
(self.serial_method << 4) | self.compression_type,
|
||||||
|
self.reserved_data,
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class Optional:
|
||||||
|
def __init__(
|
||||||
|
self, event: int = EVENT_NONE, sessionId: str = None, sequence: int = None
|
||||||
|
):
|
||||||
|
self.event = event
|
||||||
|
self.sessionId = sessionId
|
||||||
|
self.errorCode: int = 0
|
||||||
|
self.connectionId: str | None = None
|
||||||
|
self.response_meta_json: str | None = None
|
||||||
|
self.sequence = sequence
|
||||||
|
|
||||||
|
# 转成 byte 序列
|
||||||
|
def as_bytes(self) -> bytes:
|
||||||
|
option_bytes = bytearray()
|
||||||
|
if self.event != EVENT_NONE:
|
||||||
|
option_bytes.extend(self.event.to_bytes(4, "big", signed=True))
|
||||||
|
if self.sessionId is not None:
|
||||||
|
session_id_bytes = str.encode(self.sessionId)
|
||||||
|
size = len(session_id_bytes).to_bytes(4, "big", signed=True)
|
||||||
|
option_bytes.extend(size)
|
||||||
|
option_bytes.extend(session_id_bytes)
|
||||||
|
if self.sequence is not None:
|
||||||
|
option_bytes.extend(self.sequence.to_bytes(4, "big", signed=True))
|
||||||
|
return option_bytes
|
||||||
|
|
||||||
|
|
||||||
|
class Response:
|
||||||
|
def __init__(self, header: Header, optional: Optional):
|
||||||
|
self.optional = optional
|
||||||
|
self.header = header
|
||||||
|
self.payload: bytes | None = None
|
||||||
|
|
||||||
|
def __str__(self):
|
||||||
|
return super().__str__()
|
||||||
|
|
||||||
|
|
||||||
|
class TTSProvider(TTSProviderBase):
|
||||||
|
def __init__(self, config, delete_audio_file):
|
||||||
|
super().__init__(config, delete_audio_file)
|
||||||
|
self.appId = config.get("appid")
|
||||||
|
self.access_token = config.get("access_token")
|
||||||
|
self.cluster = config.get("cluster")
|
||||||
|
self.resource_id = config.get("resource_id")
|
||||||
|
self.voice = config.get("voice")
|
||||||
|
self.ws_url = config.get("ws_url")
|
||||||
|
self.authorization = config.get("authorization")
|
||||||
|
self.speaker = config.get("speaker")
|
||||||
|
self.header = {"Authorization": f"{self.authorization}{self.access_token}"}
|
||||||
|
self.stop_event_response = threading.Event()
|
||||||
|
self.enable_two_way = True
|
||||||
|
self.start_connection_flag = False
|
||||||
|
self.tts_text = ""
|
||||||
|
|
||||||
|
def startSession(self, conn):
|
||||||
|
self.conn = conn
|
||||||
|
self.tts_timeout = conn.config.get("tts_timeout", 10)
|
||||||
|
# tts 消化线程
|
||||||
|
self.tts_priority_thread = threading.Thread(
|
||||||
|
target=self._tts_priority_thread, daemon=True
|
||||||
|
)
|
||||||
|
self.tts_priority_thread.start()
|
||||||
|
|
||||||
|
# 音频播放 消化线程
|
||||||
|
self.audio_play_priority_thread = threading.Thread(
|
||||||
|
target=self._audio_play_priority_thread, daemon=True
|
||||||
|
)
|
||||||
|
self.audio_play_priority_thread.start()
|
||||||
|
ws_header = {
|
||||||
|
"X-Api-App-Key": self.appId,
|
||||||
|
"X-Api-Access-Key": self.access_token,
|
||||||
|
"X-Api-Resource-Id": self.resource_id,
|
||||||
|
"X-Api-Connect-Id": uuid.uuid4(),
|
||||||
|
}
|
||||||
|
self.ws = websockets.connect(
|
||||||
|
self.ws_url, additional_headers=ws_header, max_size=1000000000
|
||||||
|
)
|
||||||
|
tts_priority = threading.Thread(
|
||||||
|
target=self._start_monitor_tts_response_thread(), daemon=True
|
||||||
|
)
|
||||||
|
tts_priority.start()
|
||||||
|
|
||||||
|
def generate_filename(self, extension=".wav"):
|
||||||
|
return os.path.join(
|
||||||
|
self.output_file,
|
||||||
|
f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}",
|
||||||
|
)
|
||||||
|
|
||||||
|
async def send_event(
|
||||||
|
self, header: bytes, optional: bytes | None = None, payload: bytes = None
|
||||||
|
):
|
||||||
|
full_client_request = bytearray(header)
|
||||||
|
if optional is not None:
|
||||||
|
full_client_request.extend(optional)
|
||||||
|
if payload is not None:
|
||||||
|
payload_size = len(payload).to_bytes(4, "big", signed=True)
|
||||||
|
full_client_request.extend(payload_size)
|
||||||
|
full_client_request.extend(payload)
|
||||||
|
await self.ws.send(full_client_request)
|
||||||
|
|
||||||
|
async def send_text(self, speaker: str, text: str, session_id):
|
||||||
|
header = Header(
|
||||||
|
message_type=FULL_CLIENT_REQUEST,
|
||||||
|
message_type_specific_flags=MsgTypeFlagWithEvent,
|
||||||
|
serial_method=JSON,
|
||||||
|
).as_bytes()
|
||||||
|
optional = Optional(event=EVENT_TaskRequest, sessionId=session_id).as_bytes()
|
||||||
|
payload = self.get_payload_bytes(
|
||||||
|
event=EVENT_TaskRequest, text=text, speaker=speaker
|
||||||
|
)
|
||||||
|
return await self.send_event(header, optional, payload)
|
||||||
|
|
||||||
|
# 读取 res 数组某段 字符串内容
|
||||||
|
def read_res_content(self, res: bytes, offset: int):
|
||||||
|
content_size = int.from_bytes(res[offset : offset + 4], "big", signed=True)
|
||||||
|
offset += 4
|
||||||
|
content = str(res[offset : offset + content_size])
|
||||||
|
offset += content_size
|
||||||
|
return content, offset
|
||||||
|
|
||||||
|
# 读取 payload
|
||||||
|
def read_res_payload(self, res: bytes, offset: int):
|
||||||
|
payload_size = int.from_bytes(res[offset : offset + 4], "big", signed=True)
|
||||||
|
offset += 4
|
||||||
|
payload = res[offset : offset + payload_size]
|
||||||
|
offset += payload_size
|
||||||
|
return payload, offset
|
||||||
|
|
||||||
|
def parser_response(self, res) -> Response:
|
||||||
|
if isinstance(res, str):
|
||||||
|
raise RuntimeError(res)
|
||||||
|
response = Response(Header(), Optional())
|
||||||
|
# 解析结果
|
||||||
|
# header
|
||||||
|
header = response.header
|
||||||
|
num = 0b00001111
|
||||||
|
header.protocol_version = res[0] >> 4 & num
|
||||||
|
header.header_size = res[0] & 0x0F
|
||||||
|
header.message_type = (res[1] >> 4) & num
|
||||||
|
header.message_type_specific_flags = res[1] & 0x0F
|
||||||
|
header.serialization_method = res[2] >> num
|
||||||
|
header.message_compression = res[2] & 0x0F
|
||||||
|
header.reserved = res[3]
|
||||||
|
#
|
||||||
|
offset = 4
|
||||||
|
optional = response.optional
|
||||||
|
if header.message_type == FULL_SERVER_RESPONSE or AUDIO_ONLY_RESPONSE:
|
||||||
|
# read event
|
||||||
|
if header.message_type_specific_flags == MsgTypeFlagWithEvent:
|
||||||
|
optional.event = int.from_bytes(res[offset:8], "big", signed=True)
|
||||||
|
offset += 4
|
||||||
|
if optional.event == EVENT_NONE:
|
||||||
|
return response
|
||||||
|
# read connectionId
|
||||||
|
elif optional.event == EVENT_ConnectionStarted:
|
||||||
|
optional.connectionId, offset = self.read_res_content(res, offset)
|
||||||
|
elif optional.event == EVENT_ConnectionFailed:
|
||||||
|
optional.response_meta_json, offset = self.read_res_content(
|
||||||
|
res, offset
|
||||||
|
)
|
||||||
|
elif (
|
||||||
|
optional.event == EVENT_SessionStarted
|
||||||
|
or optional.event == EVENT_SessionFailed
|
||||||
|
or optional.event == EVENT_SessionFinished
|
||||||
|
):
|
||||||
|
optional.sessionId, offset = self.read_res_content(res, offset)
|
||||||
|
optional.response_meta_json, offset = self.read_res_content(
|
||||||
|
res, offset
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
optional.sessionId, offset = self.read_res_content(res, offset)
|
||||||
|
response.payload, offset = self.read_res_payload(res, offset)
|
||||||
|
|
||||||
|
elif header.message_type == ERROR_INFORMATION:
|
||||||
|
optional.errorCode = int.from_bytes(
|
||||||
|
res[offset : offset + 4], "big", signed=True
|
||||||
|
)
|
||||||
|
offset += 4
|
||||||
|
response.payload, offset = self.read_res_payload(res, offset)
|
||||||
|
return response
|
||||||
|
|
||||||
|
async def start_connection(self):
|
||||||
|
header = Header(
|
||||||
|
message_type=FULL_CLIENT_REQUEST,
|
||||||
|
message_type_specific_flags=MsgTypeFlagWithEvent,
|
||||||
|
).as_bytes()
|
||||||
|
optional = Optional(event=EVENT_Start_Connection).as_bytes()
|
||||||
|
payload = str.encode("{}")
|
||||||
|
return await self.send_event(header, optional, payload)
|
||||||
|
|
||||||
|
def print_response(self, res, tag_msg: str):
|
||||||
|
logger.bind(tag=TAG).info(f"===>{tag_msg} header:{res.header.__dict__}")
|
||||||
|
logger.bind(tag=TAG).info(f"===>{tag_msg} optional:{res.optional.__dict__}")
|
||||||
|
|
||||||
|
def get_payload_bytes(
|
||||||
|
self,
|
||||||
|
uid="1234",
|
||||||
|
event=EVENT_NONE,
|
||||||
|
text="",
|
||||||
|
speaker="",
|
||||||
|
audio_format="pcm",
|
||||||
|
audio_sample_rate=16000,
|
||||||
|
):
|
||||||
|
return str.encode(
|
||||||
|
json.dumps(
|
||||||
|
{
|
||||||
|
"user": {"uid": uid},
|
||||||
|
"event": event,
|
||||||
|
"namespace": "BidirectionalTTS",
|
||||||
|
"req_params": {
|
||||||
|
"text": text,
|
||||||
|
"speaker": speaker,
|
||||||
|
"audio_params": {
|
||||||
|
"format": audio_format,
|
||||||
|
"sample_rate": audio_sample_rate,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
async def finish_connection(self):
|
||||||
|
header = Header(
|
||||||
|
message_type=FULL_CLIENT_REQUEST,
|
||||||
|
message_type_specific_flags=MsgTypeFlagWithEvent,
|
||||||
|
serial_method=JSON,
|
||||||
|
).as_bytes()
|
||||||
|
optional = Optional(event=EVENT_FinishConnection).as_bytes()
|
||||||
|
payload = str.encode("{}")
|
||||||
|
await self.send_event(header, optional, payload)
|
||||||
|
return
|
||||||
|
|
||||||
|
async def start_session(self, session_id):
|
||||||
|
self.stop_event_response.clear()
|
||||||
|
header = Header(
|
||||||
|
message_type=FULL_CLIENT_REQUEST,
|
||||||
|
message_type_specific_flags=MsgTypeFlagWithEvent,
|
||||||
|
serial_method=JSON,
|
||||||
|
).as_bytes()
|
||||||
|
optional = Optional(event=EVENT_StartSession, sessionId=session_id).as_bytes()
|
||||||
|
payload = self.get_payload_bytes(event=EVENT_StartSession, speaker=self.speaker)
|
||||||
|
await self.send_event(header, optional, payload)
|
||||||
|
|
||||||
|
async def finish_session(self, session_id):
|
||||||
|
self.stop_event_response.set()
|
||||||
|
header = Header(
|
||||||
|
message_type=FULL_CLIENT_REQUEST,
|
||||||
|
message_type_specific_flags=MsgTypeFlagWithEvent,
|
||||||
|
serial_method=JSON,
|
||||||
|
).as_bytes()
|
||||||
|
optional = Optional(event=EVENT_FinishSession, sessionId=session_id).as_bytes()
|
||||||
|
payload = str.encode("{}")
|
||||||
|
await self.send_event(header, optional, payload)
|
||||||
|
return
|
||||||
|
|
||||||
|
async def reset(self):
|
||||||
|
# 关闭之前的对话
|
||||||
|
if self.start_connection_flag:
|
||||||
|
await self.finish_connection()
|
||||||
|
self.start_connection_flag = False
|
||||||
|
await self.start_connection()
|
||||||
|
self.start_connection_flag = True
|
||||||
|
await super().reset()
|
||||||
|
|
||||||
|
async def close(self):
|
||||||
|
super().close()
|
||||||
|
"""资源清理方法"""
|
||||||
|
await self.finish_connection()
|
||||||
|
await self.ws.close()
|
||||||
|
|
||||||
|
async def text_to_speak(self, text, _):
|
||||||
|
# 发送文本
|
||||||
|
await self.send_text(self.speaker, text, self.conn.session_id)
|
||||||
|
return
|
||||||
|
|
||||||
|
def _start_monitor_tts_response_thread(self):
|
||||||
|
# 初始化链接
|
||||||
|
asyncio.run_coroutine_threadsafe(
|
||||||
|
self._start_monitor_tts_response(), loop=self.conn.loop
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _start_monitor_tts_response(self):
|
||||||
|
chunk_total = b""
|
||||||
|
while not self.stop_event.is_set():
|
||||||
|
try:
|
||||||
|
msg = await self.ws.recv() # 确保 `recv()` 运行在同一个 event loop
|
||||||
|
res = self.parser_response(msg)
|
||||||
|
self.print_response(res, "send_text res:")
|
||||||
|
|
||||||
|
if (
|
||||||
|
res.optional.event == EVENT_TTSResponse
|
||||||
|
and res.header.message_type == AUDIO_ONLY_RESPONSE
|
||||||
|
):
|
||||||
|
logger.bind(tag=TAG).info(f"推送数据到队列里面~~")
|
||||||
|
opus_datas = self.wav_to_opus_data_audio_raw(res.payload)
|
||||||
|
logger.bind(tag=TAG).info(
|
||||||
|
f"推送数据到队列里面帧数~~{len(opus_datas)}"
|
||||||
|
)
|
||||||
|
self.audio_play_queue.put((opus_datas, None, 0))
|
||||||
|
elif res.optional.event == EVENT_TTSSentenceStart:
|
||||||
|
json_data = json.loads(res.payload.decode("utf-8"))
|
||||||
|
self.tts_text = json_data.get("text", "")
|
||||||
|
logger.bind(tag=TAG).info(f"句子开始~~{self.tts_text}")
|
||||||
|
self.audio_play_queue.put((None, self.tts_text, 0))
|
||||||
|
|
||||||
|
elif res.optional.event == EVENT_TTSSentenceEnd:
|
||||||
|
logger.bind(tag=TAG).info(f"句子结束~~{self.tts_text}")
|
||||||
|
self.audio_play_queue.put((None, self.tts_text, 0))
|
||||||
|
elif res.optional.event == EVENT_SessionFinished:
|
||||||
|
logger.bind(tag=TAG).info(f"会话结束~~,最后一句补零")
|
||||||
|
# opus_datas = self.wav_to_opus_data_audio_raw(b"", is_end=True)
|
||||||
|
self.audio_play_queue.put((None, self.tts_text, 0))
|
||||||
|
else:
|
||||||
|
continue
|
||||||
|
except websockets.ConnectionClosed:
|
||||||
|
break # 连接关闭时退出监听
|
||||||
|
except Exception as e:
|
||||||
|
logger.bind(tag=TAG).error(f"Error in _start_monitor_tts_response: {e}")
|
||||||
|
traceback.print_exc()
|
||||||
|
continue
|
||||||
@@ -0,0 +1,132 @@
|
|||||||
|
import array
|
||||||
|
import logging
|
||||||
|
import traceback
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
from typing import List, Optional
|
||||||
|
from opuslib_next import Encoder
|
||||||
|
from opuslib_next import constants
|
||||||
|
|
||||||
|
|
||||||
|
class OpusEncoderUtils:
|
||||||
|
"""PCM到Opus的编码器"""
|
||||||
|
|
||||||
|
def __init__(self, sample_rate: int, channels: int, frame_size_ms: int):
|
||||||
|
"""
|
||||||
|
初始化Opus编码器
|
||||||
|
|
||||||
|
Args:
|
||||||
|
sample_rate: 采样率 (Hz)
|
||||||
|
channels: 通道数 (1=单声道, 2=立体声)
|
||||||
|
frame_size_ms: 帧大小 (毫秒)
|
||||||
|
"""
|
||||||
|
self.sample_rate = sample_rate
|
||||||
|
self.channels = channels
|
||||||
|
self.frame_size_ms = frame_size_ms
|
||||||
|
# 计算每帧样本数 = 采样率 * 帧大小(毫秒) / 1000
|
||||||
|
self.frame_size = (sample_rate * frame_size_ms) // 1000
|
||||||
|
# 总帧大小 = 每帧样本数 * 通道数
|
||||||
|
self.total_frame_size = self.frame_size * channels
|
||||||
|
|
||||||
|
# 比特率和复杂度设置
|
||||||
|
self.bitrate = 24000 # bps
|
||||||
|
self.complexity = 10 # 最高质量
|
||||||
|
|
||||||
|
# 缓冲区初始化为空
|
||||||
|
self.buffer = np.array([], dtype=np.int16)
|
||||||
|
|
||||||
|
try:
|
||||||
|
# 创建Opus编码器
|
||||||
|
self.encoder = Encoder(
|
||||||
|
sample_rate, channels, constants.APPLICATION_AUDIO # 音频优化模式
|
||||||
|
)
|
||||||
|
self.encoder.bitrate = self.bitrate
|
||||||
|
self.encoder.complexity = self.complexity
|
||||||
|
self.encoder.signal = constants.SIGNAL_VOICE # 语音信号优化
|
||||||
|
except Exception as e:
|
||||||
|
logging.error(f"初始化Opus编码器失败: {e}")
|
||||||
|
raise RuntimeError("初始化失败") from e
|
||||||
|
|
||||||
|
def reset_state(self):
|
||||||
|
"""重置编码器状态"""
|
||||||
|
self.encoder.reset_state()
|
||||||
|
self.buffer = np.array([], dtype=np.int16)
|
||||||
|
|
||||||
|
def encode_pcm_to_opus(self, pcm_data: bytes, end_of_stream: bool) -> List[bytes]:
|
||||||
|
"""
|
||||||
|
将PCM数据编码为Opus格式
|
||||||
|
|
||||||
|
Args:
|
||||||
|
pcm_data: PCM字节数据
|
||||||
|
end_of_stream: 是否为流的结束
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Opus数据包列表
|
||||||
|
"""
|
||||||
|
# 将字节数据转换为short数组
|
||||||
|
new_samples = self._convert_bytes_to_shorts(pcm_data)
|
||||||
|
|
||||||
|
# 校验PCM数据
|
||||||
|
self._validate_pcm_data(new_samples)
|
||||||
|
|
||||||
|
# 将新数据追加到缓冲区
|
||||||
|
self.buffer = np.append(self.buffer, new_samples)
|
||||||
|
|
||||||
|
opus_packets = []
|
||||||
|
offset = 0
|
||||||
|
|
||||||
|
# 处理所有完整帧
|
||||||
|
while offset <= len(self.buffer) - self.total_frame_size:
|
||||||
|
frame = self.buffer[offset : offset + self.total_frame_size]
|
||||||
|
output = self._encode(frame)
|
||||||
|
if output:
|
||||||
|
opus_packets.append(output)
|
||||||
|
offset += self.total_frame_size
|
||||||
|
|
||||||
|
# 保留未处理的样本
|
||||||
|
self.buffer = self.buffer[offset:]
|
||||||
|
|
||||||
|
# 流结束时处理剩余数据
|
||||||
|
if end_of_stream and len(self.buffer) > 0:
|
||||||
|
# 创建最后一帧并用0填充
|
||||||
|
last_frame = np.zeros(self.total_frame_size, dtype=np.int16)
|
||||||
|
last_frame[: len(self.buffer)] = self.buffer
|
||||||
|
|
||||||
|
output = self._encode(last_frame)
|
||||||
|
if output:
|
||||||
|
opus_packets.append(output)
|
||||||
|
self.buffer = np.array([], dtype=np.int16)
|
||||||
|
|
||||||
|
return opus_packets
|
||||||
|
|
||||||
|
def _encode(self, frame: np.ndarray) -> Optional[bytes]:
|
||||||
|
"""编码一帧音频数据"""
|
||||||
|
try:
|
||||||
|
# 将numpy数组转换为bytes
|
||||||
|
frame_bytes = frame.tobytes()
|
||||||
|
# opuslib要求输入字节数必须是channels*2的倍数
|
||||||
|
encoded = self.encoder.encode(frame_bytes, self.frame_size)
|
||||||
|
return encoded
|
||||||
|
except Exception as e:
|
||||||
|
logging.error(f"Opus编码失败: {e}")
|
||||||
|
traceback.print_exc()
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _convert_bytes_to_shorts(self, bytes_data: bytes) -> np.ndarray:
|
||||||
|
"""将字节数组转换为short数组 (16位PCM)"""
|
||||||
|
# 假设输入是小端字节序的16位PCM
|
||||||
|
return np.frombuffer(bytes_data, dtype=np.int16)
|
||||||
|
|
||||||
|
def _validate_pcm_data(self, pcm_shorts: np.ndarray) -> None:
|
||||||
|
"""验证PCM数据是否有效"""
|
||||||
|
# 16位PCM数据范围是 -32768 到 32767
|
||||||
|
if np.any((pcm_shorts < -32768) | (pcm_shorts > 32767)):
|
||||||
|
invalid_samples = pcm_shorts[(pcm_shorts < -32768) | (pcm_shorts > 32767)]
|
||||||
|
logging.warning(f"发现无效PCM样本: {invalid_samples[:5]}...")
|
||||||
|
# 在实际应用中可以选择裁剪而不是抛出异常
|
||||||
|
# np.clip(pcm_shorts, -32768, 32767, out=pcm_shorts)
|
||||||
|
|
||||||
|
def close(self):
|
||||||
|
"""关闭编码器并释放资源"""
|
||||||
|
# opuslib没有明确的关闭方法,Python的垃圾回收会处理
|
||||||
|
pass
|
||||||
Reference in New Issue
Block a user