feat: 开放小爱音箱接入小智 AI 演示源代码
This commit is contained in:
@@ -0,0 +1,271 @@
|
||||
import logging
|
||||
import queue
|
||||
import numpy as np
|
||||
import pyaudio
|
||||
import opuslib
|
||||
from xiaozhi.services.protocols.typing import AudioConfig
|
||||
import time
|
||||
import sys
|
||||
|
||||
from xiaozhi.services.audio.stream import MyAudio
|
||||
from xiaozhi.xiaoai import XiaoAi
|
||||
|
||||
logger = logging.getLogger("AudioCodec")
|
||||
|
||||
|
||||
class AudioCodec:
|
||||
"""音频编解码器类,处理音频的录制和播放"""
|
||||
|
||||
def __init__(self):
|
||||
"""初始化音频编解码器"""
|
||||
self.audio = None
|
||||
self.input_stream = None
|
||||
self.output_stream = None
|
||||
self.opus_encoder = None
|
||||
self.opus_decoder = None
|
||||
self.audio_decode_queue = queue.Queue()
|
||||
self._is_closing = False
|
||||
|
||||
self._initialize_audio()
|
||||
|
||||
def _initialize_audio(self):
|
||||
"""初始化音频设备和编解码器"""
|
||||
try:
|
||||
self.audio = MyAudio() if XiaoAi.mode == "xiaoai" else pyaudio.PyAudio()
|
||||
|
||||
# 初始化音频输入流
|
||||
self.input_stream = self.audio.open(
|
||||
format=pyaudio.paInt16,
|
||||
channels=AudioConfig.CHANNELS,
|
||||
rate=AudioConfig.SAMPLE_RATE,
|
||||
input=True,
|
||||
frames_per_buffer=AudioConfig.FRAME_SIZE,
|
||||
)
|
||||
|
||||
# 初始化音频输出流
|
||||
self.output_stream = self.audio.open(
|
||||
format=pyaudio.paInt16,
|
||||
channels=AudioConfig.CHANNELS,
|
||||
rate=AudioConfig.SAMPLE_RATE,
|
||||
output=True,
|
||||
frames_per_buffer=AudioConfig.FRAME_SIZE,
|
||||
)
|
||||
|
||||
# 初始化Opus编码器
|
||||
self.opus_encoder = opuslib.Encoder(
|
||||
fs=AudioConfig.SAMPLE_RATE,
|
||||
channels=AudioConfig.CHANNELS,
|
||||
application=opuslib.APPLICATION_AUDIO,
|
||||
)
|
||||
|
||||
# 初始化Opus解码器
|
||||
self.opus_decoder = opuslib.Decoder(
|
||||
fs=AudioConfig.SAMPLE_RATE, channels=AudioConfig.CHANNELS
|
||||
)
|
||||
|
||||
logger.info("音频设备和编解码器初始化成功")
|
||||
except Exception as e:
|
||||
logger.error(f"初始化音频设备失败: {e}")
|
||||
raise
|
||||
|
||||
def read_audio(self):
|
||||
"""读取音频输入数据并编码"""
|
||||
try:
|
||||
data = self.input_stream.read(
|
||||
AudioConfig.FRAME_SIZE, exception_on_overflow=False
|
||||
)
|
||||
if not data:
|
||||
return None
|
||||
return self.opus_encoder.encode(data, AudioConfig.FRAME_SIZE)
|
||||
except Exception as e:
|
||||
logger.error(f"读取音频输入时出错: {e}")
|
||||
return None
|
||||
|
||||
def write_audio(self, opus_data):
|
||||
"""将编码的音频数据添加到播放队列"""
|
||||
self.audio_decode_queue.put(opus_data)
|
||||
|
||||
def play_audio(self):
|
||||
"""处理并播放队列中的音频数据"""
|
||||
try:
|
||||
# 批量处理多个音频包以减少处理延迟
|
||||
batch_size = min(10, self.audio_decode_queue.qsize())
|
||||
if batch_size == 0:
|
||||
return False
|
||||
|
||||
# 创建缓冲区存储解码后的数据
|
||||
buffer = bytearray()
|
||||
|
||||
for _ in range(batch_size):
|
||||
if self.audio_decode_queue.empty():
|
||||
break
|
||||
|
||||
opus_data = self.audio_decode_queue.get_nowait()
|
||||
try:
|
||||
pcm_data = self.opus_decoder.decode(
|
||||
opus_data, AudioConfig.FRAME_SIZE, decode_fec=False
|
||||
)
|
||||
buffer.extend(pcm_data)
|
||||
except Exception as e:
|
||||
logger.error(f"解码音频数据时出错: {e}")
|
||||
|
||||
# 只有在有数据时才处理和播放
|
||||
if len(buffer) > 0:
|
||||
# 转换为numpy数组
|
||||
pcm_array = np.frombuffer(buffer, dtype=np.int16)
|
||||
|
||||
# 播放音频
|
||||
try:
|
||||
if self.output_stream and self.output_stream.is_active():
|
||||
self.output_stream.write(pcm_array.tobytes())
|
||||
return True
|
||||
else:
|
||||
# MAC 特定:如果流不活跃,尝试重新初始化
|
||||
self._reinitialize_output_stream()
|
||||
if self.output_stream and self.output_stream.is_active():
|
||||
self.output_stream.write(pcm_array.tobytes())
|
||||
return True
|
||||
except OSError as e:
|
||||
if "Stream closed" in str(e) or "Internal PortAudio error" in str(
|
||||
e
|
||||
):
|
||||
logger.error(f"播放音频时出错: {e}")
|
||||
self._reinitialize_output_stream()
|
||||
else:
|
||||
logger.error(f"播放音频时出错: {e}")
|
||||
except queue.Empty:
|
||||
pass
|
||||
except Exception as e:
|
||||
logger.error(f"播放音频时出错: {e}")
|
||||
self._reinitialize_output_stream()
|
||||
|
||||
return False
|
||||
|
||||
def has_pending_audio(self):
|
||||
"""检查是否还有待播放的音频数据"""
|
||||
return not self.audio_decode_queue.empty()
|
||||
|
||||
def wait_for_audio_complete(self, timeout=5.0):
|
||||
# 等待音频队列清空
|
||||
attempt = 0
|
||||
max_attempts = 15
|
||||
while not self.audio_decode_queue.empty() and attempt < max_attempts:
|
||||
time.sleep(0.1)
|
||||
attempt += 1
|
||||
|
||||
# 在关闭前清空任何剩余数据
|
||||
while not self.audio_decode_queue.empty():
|
||||
try:
|
||||
self.audio_decode_queue.get_nowait()
|
||||
except queue.Empty:
|
||||
break
|
||||
|
||||
def clear_audio_queue(self):
|
||||
"""清空音频队列"""
|
||||
while not self.audio_decode_queue.empty():
|
||||
try:
|
||||
self.audio_decode_queue.get_nowait()
|
||||
except queue.Empty:
|
||||
break
|
||||
|
||||
def start_streams(self):
|
||||
"""启动音频流"""
|
||||
if not self.input_stream.is_active():
|
||||
self.input_stream.start_stream()
|
||||
if not self.output_stream.is_active():
|
||||
self.output_stream.start_stream()
|
||||
|
||||
def stop_streams(self):
|
||||
"""停止音频流"""
|
||||
if self.input_stream and self.input_stream.is_active():
|
||||
self.input_stream.stop_stream()
|
||||
if self.output_stream and self.output_stream.is_active():
|
||||
self.output_stream.stop_stream()
|
||||
|
||||
def _reinitialize_output_stream(self):
|
||||
"""重新初始化音频输出流"""
|
||||
if self._is_closing: # 如果正在关闭,不要重新初始化
|
||||
return
|
||||
|
||||
try:
|
||||
if self.output_stream:
|
||||
try:
|
||||
if self.output_stream.is_active():
|
||||
self.output_stream.stop_stream()
|
||||
self.output_stream.close()
|
||||
except Exception as e:
|
||||
# logger.warning(f"关闭旧输出流时出错: {e}")
|
||||
pass
|
||||
|
||||
# 在 MAC 上添加短暂延迟
|
||||
if sys.platform in ("darwin", "linux"):
|
||||
time.sleep(0.1)
|
||||
|
||||
self.output_stream = self.audio.open(
|
||||
format=pyaudio.paInt16,
|
||||
channels=AudioConfig.CHANNELS,
|
||||
rate=AudioConfig.SAMPLE_RATE,
|
||||
output=True,
|
||||
frames_per_buffer=AudioConfig.FRAME_SIZE,
|
||||
)
|
||||
logger.info("音频输出流重新初始化成功")
|
||||
except Exception as e:
|
||||
logger.error(f"重新初始化音频输出流失败: {e}")
|
||||
raise
|
||||
|
||||
def close(self):
|
||||
"""关闭音频编解码器,确保资源正确释放"""
|
||||
if self._is_closing: # 防止重复关闭
|
||||
return
|
||||
|
||||
self._is_closing = True
|
||||
logger.info("开始关闭音频编解码器...")
|
||||
|
||||
try:
|
||||
# 等待并清理剩余音频数据
|
||||
self.wait_for_audio_complete()
|
||||
|
||||
# 关闭输入流
|
||||
if self.input_stream:
|
||||
logger.debug("正在关闭输入流...")
|
||||
try:
|
||||
if self.input_stream.is_active():
|
||||
self.input_stream.stop_stream()
|
||||
self.input_stream.close()
|
||||
except Exception as e:
|
||||
logger.error(f"关闭输入流时出错: {e}")
|
||||
self.input_stream = None
|
||||
|
||||
# 关闭输出流
|
||||
if self.output_stream:
|
||||
logger.debug("正在关闭输出流...")
|
||||
try:
|
||||
if self.output_stream.is_active():
|
||||
self.output_stream.stop_stream()
|
||||
self.output_stream.close()
|
||||
except Exception as e:
|
||||
logger.error(f"关闭输出流时出错: {e}")
|
||||
self.output_stream = None
|
||||
|
||||
# 关闭 PyAudio 实例
|
||||
if self.audio:
|
||||
logger.debug("正在终止 PyAudio...")
|
||||
try:
|
||||
self.audio.terminate()
|
||||
except Exception as e:
|
||||
logger.error(f"终止 PyAudio 时出错: {e}")
|
||||
self.audio = None
|
||||
|
||||
# 清理编解码器
|
||||
self.opus_encoder = None
|
||||
self.opus_decoder = None
|
||||
|
||||
logger.info("音频编解码器关闭完成")
|
||||
except Exception as e:
|
||||
logger.error(f"关闭音频编解码器时发生错误: {e}")
|
||||
finally:
|
||||
self._is_closing = False
|
||||
|
||||
def __del__(self):
|
||||
"""析构函数,确保资源被释放"""
|
||||
self.close()
|
||||
Binary file not shown.
@@ -0,0 +1,45 @@
|
||||
import os
|
||||
import sys
|
||||
import subprocess
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger("Opus")
|
||||
|
||||
|
||||
def setup_opus():
|
||||
libs_dir = ""
|
||||
lib_path = ""
|
||||
|
||||
logger.info("正在加载 Opus 动态库,请耐心等待...")
|
||||
|
||||
if sys.platform == "win32":
|
||||
libs_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
lib_path = os.path.join(libs_dir, "opus.dll")
|
||||
elif sys.platform == "darwin":
|
||||
result = subprocess.check_output(
|
||||
"brew list opus | grep libopus.dylib", shell=True
|
||||
)
|
||||
lib_path = result.decode("utf-8").strip()
|
||||
if not lib_path.endswith("libopus.dylib"):
|
||||
raise RuntimeError("请先安装 Opus: brew install opus")
|
||||
libs_dir = os.path.dirname(lib_path)
|
||||
else:
|
||||
raise RuntimeError(f"暂不支持 {sys.platform} 平台")
|
||||
|
||||
import ctypes.util
|
||||
|
||||
original_find_library = ctypes.util.find_library
|
||||
|
||||
def patched_find_library(name):
|
||||
if name == "opus":
|
||||
return lib_path
|
||||
return original_find_library(name)
|
||||
|
||||
ctypes.util.find_library = patched_find_library
|
||||
|
||||
if hasattr(os, "add_dll_directory"):
|
||||
os.add_dll_directory(libs_dir)
|
||||
|
||||
os.environ["PATH"] = libs_dir + os.pathsep + os.environ.get("PATH", "")
|
||||
|
||||
ctypes.CDLL(lib_path)
|
||||
@@ -0,0 +1,290 @@
|
||||
from typing import ClassVar, Optional, Callable, Any
|
||||
from threading import Lock
|
||||
from collections import deque
|
||||
|
||||
|
||||
class GlobalStream:
|
||||
"""
|
||||
全局音频缓冲区,用于存储和分发音频数据
|
||||
|
||||
用来将小爱音箱的音频输入输出流,适配到小智 AI 的音频输入输出流
|
||||
"""
|
||||
|
||||
_instance = None
|
||||
_lock = Lock()
|
||||
|
||||
# 默认缓冲区大小(以字节为单位)
|
||||
DEFAULT_BUFFER_SIZE = 1024 * 1024 # 1MB
|
||||
|
||||
def __new__(cls):
|
||||
with cls._lock:
|
||||
if cls._instance is None:
|
||||
cls._instance = super(GlobalStream, cls).__new__(cls)
|
||||
cls._instance._initialize()
|
||||
return cls._instance
|
||||
|
||||
def _initialize(self):
|
||||
"""初始化实例变量"""
|
||||
self._max_buffer_size = self.DEFAULT_BUFFER_SIZE
|
||||
self._input_buffer = deque(maxlen=self._max_buffer_size)
|
||||
self._is_input_active = False
|
||||
self._is_output_active = False
|
||||
self._input_readers = {} # 跟踪每个流的读取位置
|
||||
self._reader_counter = 0 # 为每个读取器分配唯一ID
|
||||
self._buffer_overflow_count = 0 # 记录缓冲区溢出次数
|
||||
self.on_output_data = None # 输入数据回调函数
|
||||
|
||||
def set_buffer_size(self, frames: int) -> None:
|
||||
if frames <= 0:
|
||||
raise ValueError("缓冲区大小必须大于0")
|
||||
|
||||
# 创建新的有限长度缓冲区
|
||||
new_input_buffer = deque(self._input_buffer, maxlen=frames * 2)
|
||||
|
||||
# 替换旧缓冲区
|
||||
self._input_buffer = new_input_buffer
|
||||
self._max_buffer_size = frames * 2
|
||||
|
||||
# 重置读取位置,因为缓冲区可能已经改变
|
||||
for reader_id in self._input_readers:
|
||||
self._input_readers[reader_id] = 0
|
||||
|
||||
def register_reader(self) -> int:
|
||||
"""注册一个新的读取器并返回其ID"""
|
||||
reader_id = self._reader_counter
|
||||
self._input_readers[reader_id] = 0 # 初始位置为0
|
||||
self._reader_counter += 1
|
||||
return reader_id
|
||||
|
||||
def unregister_reader(self, reader_id: int) -> None:
|
||||
"""注销一个读取器"""
|
||||
if reader_id in self._input_readers:
|
||||
del self._input_readers[reader_id]
|
||||
|
||||
def read(self, reader_id: int, num_frames: int) -> bytes:
|
||||
num_frames = num_frames * 2
|
||||
if not self._is_input_active:
|
||||
return bytes(num_frames)
|
||||
|
||||
if reader_id not in self._input_readers:
|
||||
return bytes(num_frames)
|
||||
|
||||
# 将输入缓冲区转换为列表以便随机访问
|
||||
buffer_list = list(self._input_buffer)
|
||||
current_pos = self._input_readers[reader_id]
|
||||
|
||||
# 如果当前位置超出缓冲区大小,返回空字节
|
||||
if current_pos >= len(buffer_list):
|
||||
return bytes(num_frames)
|
||||
|
||||
# 读取数据
|
||||
end_pos = min(current_pos + num_frames, len(buffer_list))
|
||||
data = bytes(buffer_list[current_pos:end_pos])
|
||||
|
||||
# 更新读取位置
|
||||
self._input_readers[reader_id] = end_pos
|
||||
|
||||
# 如果数据不足,用零填充
|
||||
if len(data) < num_frames:
|
||||
data += bytes(num_frames - len(data))
|
||||
|
||||
return data
|
||||
|
||||
# 转发输出音频流
|
||||
def write(self, frames: bytes) -> None:
|
||||
"""写入数据到输出缓冲区"""
|
||||
if not self._is_output_active:
|
||||
return
|
||||
|
||||
if self.on_output_data:
|
||||
self.on_output_data(frames)
|
||||
|
||||
# 添加输入音频流
|
||||
def add_input_data(self, data: bytes) -> None:
|
||||
if not self._is_input_active:
|
||||
return
|
||||
|
||||
"""添加输入数据到全局缓冲区"""
|
||||
# 检查是否会溢出
|
||||
if len(self._input_buffer) + len(data) > self._max_buffer_size:
|
||||
self._buffer_overflow_count += 1
|
||||
|
||||
for b in data:
|
||||
self._input_buffer.append(b)
|
||||
|
||||
# 如果有读取器的位置已经超出了缓冲区大小,需要调整
|
||||
if len(self._input_buffer) >= self._max_buffer_size:
|
||||
for reader_id in self._input_readers:
|
||||
if self._input_readers[reader_id] > len(self._input_buffer) // 2:
|
||||
# 将读取位置重置到缓冲区中间,避免读取器永远跟不上
|
||||
self._input_readers[reader_id] = len(self._input_buffer) // 2
|
||||
|
||||
def start_input(self) -> None:
|
||||
"""启动输入流"""
|
||||
self._is_input_active = True
|
||||
|
||||
def stop_input(self) -> None:
|
||||
"""停止输入流"""
|
||||
self._is_input_active = False
|
||||
self.clear()
|
||||
|
||||
def start_output(self) -> None:
|
||||
"""启动输出流"""
|
||||
self._is_output_active = True
|
||||
|
||||
def stop_output(self) -> None:
|
||||
"""停止输出流"""
|
||||
self._is_output_active = False
|
||||
|
||||
def is_input_active(self) -> bool:
|
||||
"""检查输入流是否活跃"""
|
||||
return self._is_input_active
|
||||
|
||||
def is_output_active(self) -> bool:
|
||||
"""检查输出流是否活跃"""
|
||||
return self._is_output_active
|
||||
|
||||
def clear(self) -> None:
|
||||
"""清空缓冲区"""
|
||||
self._input_buffer.clear()
|
||||
for reader_id in self._input_readers:
|
||||
self._input_readers[reader_id] = 0
|
||||
|
||||
|
||||
class MyStream:
|
||||
"""音频流类,用于读写音频数据"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
rate: int,
|
||||
channels: int,
|
||||
format: int,
|
||||
input: bool = False,
|
||||
output: bool = False,
|
||||
frames_per_buffer: int = 1024,
|
||||
start: bool = True,
|
||||
) -> None:
|
||||
self._rate = rate
|
||||
self._channels = channels
|
||||
self._format = format
|
||||
self._frames_per_buffer = frames_per_buffer
|
||||
self._is_input = input
|
||||
self._is_output = output
|
||||
self._is_active = False
|
||||
self._is_closed = False
|
||||
|
||||
# 获取全局音频缓冲区
|
||||
self._global_buffer = GlobalStream()
|
||||
|
||||
# 注册读取器ID
|
||||
self._reader_id = self._global_buffer.register_reader() if input else None
|
||||
|
||||
# 如果需要,启动流
|
||||
if start:
|
||||
self.start_stream()
|
||||
|
||||
def close(self) -> None:
|
||||
"""关闭流"""
|
||||
if not self._is_closed:
|
||||
self.stop_stream()
|
||||
if self._is_input and self._reader_id is not None:
|
||||
self._global_buffer.unregister_reader(self._reader_id)
|
||||
self._is_closed = True
|
||||
|
||||
def is_active(self) -> bool:
|
||||
"""检查流是否活跃"""
|
||||
return self._is_active
|
||||
|
||||
def start_stream(self) -> None:
|
||||
"""启动流"""
|
||||
if not self._is_active and not self._is_closed:
|
||||
self._is_active = True
|
||||
if self._is_input:
|
||||
self._global_buffer.start_input()
|
||||
if self._is_output:
|
||||
self._global_buffer.start_output()
|
||||
|
||||
def stop_stream(self) -> None:
|
||||
"""停止流"""
|
||||
if self._is_active:
|
||||
self._is_active = False
|
||||
if self._is_input:
|
||||
self._global_buffer.stop_input()
|
||||
if self._is_output:
|
||||
self._global_buffer.stop_output()
|
||||
|
||||
def read(self, num_frames: int, exception_on_overflow=False) -> bytes:
|
||||
"""从输入流读取数据"""
|
||||
if (
|
||||
not self._is_input
|
||||
or self._is_closed
|
||||
or not self._is_active
|
||||
or self._reader_id is None
|
||||
):
|
||||
return bytes(num_frames)
|
||||
|
||||
return self._global_buffer.read(self._reader_id, num_frames)
|
||||
|
||||
def write(self, frames: bytes) -> None:
|
||||
"""写入数据到输出流"""
|
||||
if not self._is_output or self._is_closed or not self._is_active:
|
||||
return
|
||||
|
||||
self._global_buffer.write(frames)
|
||||
|
||||
|
||||
class MyAudio:
|
||||
"""PyAudio替代品,用于创建和管理音频流"""
|
||||
|
||||
Stream: ClassVar[type] = MyStream
|
||||
|
||||
def __init__(self, buffer_size: int = GlobalStream.DEFAULT_BUFFER_SIZE) -> None:
|
||||
# 初始化全局音频缓冲区
|
||||
self._global_buffer = GlobalStream()
|
||||
# 设置缓冲区大小
|
||||
if buffer_size != GlobalStream.DEFAULT_BUFFER_SIZE:
|
||||
self._global_buffer.set_buffer_size(buffer_size)
|
||||
self._is_terminated = False
|
||||
|
||||
def open(
|
||||
self,
|
||||
rate: int,
|
||||
channels: int,
|
||||
format: int,
|
||||
input: bool = False,
|
||||
output: bool = False,
|
||||
input_device_index: Optional[int] = None,
|
||||
output_device_index: Optional[int] = None,
|
||||
frames_per_buffer: int = 1024,
|
||||
start: bool = True,
|
||||
input_host_api_specific_stream_info: Optional[Any] = None,
|
||||
output_host_api_specific_stream_info: Optional[Any] = None,
|
||||
stream_callback: Optional[Callable] = None,
|
||||
) -> MyStream:
|
||||
"""打开一个新的音频流"""
|
||||
if self._is_terminated:
|
||||
raise RuntimeError("MyAudio instance has been terminated")
|
||||
|
||||
# 启动全局输入/输出流(如果需要)
|
||||
if input:
|
||||
self._global_buffer.start_input()
|
||||
if output:
|
||||
self._global_buffer.start_output()
|
||||
|
||||
# 创建并返回一个新的流实例
|
||||
return MyStream(
|
||||
rate=rate,
|
||||
channels=channels,
|
||||
format=format,
|
||||
input=input,
|
||||
output=output,
|
||||
frames_per_buffer=frames_per_buffer,
|
||||
start=start,
|
||||
)
|
||||
|
||||
def terminate(self) -> None:
|
||||
"""终止MyAudio实例"""
|
||||
if not self._is_terminated:
|
||||
self._global_buffer.stop_input()
|
||||
self._global_buffer.stop_output()
|
||||
self._is_terminated = True
|
||||
@@ -0,0 +1,53 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Optional, Callable
|
||||
import logging
|
||||
|
||||
class BaseDisplay(ABC):
|
||||
"""显示接口的抽象基类"""
|
||||
|
||||
def __init__(self):
|
||||
self.logger = logging.getLogger(self.__class__.__name__)
|
||||
self.current_volume = 70 # 默认音量
|
||||
|
||||
@abstractmethod
|
||||
def set_callbacks(self,
|
||||
press_callback: Optional[Callable] = None,
|
||||
release_callback: Optional[Callable] = None,
|
||||
status_callback: Optional[Callable] = None,
|
||||
text_callback: Optional[Callable] = None,
|
||||
emotion_callback: Optional[Callable] = None,
|
||||
mode_callback: Optional[Callable] = None,
|
||||
auto_callback: Optional[Callable] = None,
|
||||
abort_callback: Optional[Callable] = None): # 添加打断回调参数
|
||||
"""设置回调函数"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def update_button_status(self, text: str):
|
||||
"""更新按钮状态"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def update_status(self, status: str):
|
||||
"""更新状态文本"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def update_text(self, text: str):
|
||||
"""更新TTS文本"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def update_emotion(self, emotion: str):
|
||||
"""更新表情"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def start(self):
|
||||
"""启动显示"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def on_close(self):
|
||||
"""关闭显示"""
|
||||
pass
|
||||
@@ -0,0 +1,284 @@
|
||||
import threading
|
||||
import tkinter as tk
|
||||
from tkinter import ttk
|
||||
import queue
|
||||
import logging
|
||||
import time
|
||||
from typing import Optional, Callable
|
||||
|
||||
from xiaozhi.services.display.base_display import BaseDisplay
|
||||
|
||||
|
||||
class GuiDisplay(BaseDisplay):
|
||||
def __init__(self):
|
||||
super().__init__() # 调用父类初始化
|
||||
"""创建 GUI 界面"""
|
||||
# 初始化日志
|
||||
self.logger = logging.getLogger("Display")
|
||||
|
||||
# 创建主窗口
|
||||
self.root = tk.Tk()
|
||||
self.root.title("小爱音箱接入小智 AI 演示")
|
||||
self.root.geometry("520x360")
|
||||
|
||||
# 在窗口底部添加作者信息
|
||||
self.author_label = ttk.Label(self.root, text="作者: https://del.wang")
|
||||
self.author_label.pack(side=tk.BOTTOM, pady=5)
|
||||
|
||||
# 让窗口居中显示
|
||||
self.root.update_idletasks()
|
||||
width = self.root.winfo_width()
|
||||
height = self.root.winfo_height()
|
||||
x = (self.root.winfo_screenwidth() // 2) - (width // 2)
|
||||
y = (self.root.winfo_screenheight() // 2) - (height // 2)
|
||||
self.root.geometry(f"+{x}+{y}")
|
||||
|
||||
# 状态显示
|
||||
self.status_frame = ttk.Frame(self.root)
|
||||
self.status_frame.pack(pady=20)
|
||||
self.status_label = ttk.Label(self.status_frame, text="状态: 未连接")
|
||||
self.status_label.pack(side=tk.LEFT)
|
||||
|
||||
# 表情显示
|
||||
self.emotion_label = tk.Label(self.root, text="😊", font=("Segoe UI Emoji", 32))
|
||||
self.emotion_label.pack(padx=20, pady=20)
|
||||
|
||||
# TTS文本显示
|
||||
self.tts_text_label = ttk.Label(
|
||||
self.root, text="很高兴认识你!", wraplength=250
|
||||
)
|
||||
self.tts_text_label.pack(padx=20, pady=10)
|
||||
|
||||
# 控制按钮
|
||||
self.btn_frame = ttk.Frame(self.root)
|
||||
self.btn_frame.pack(pady=20)
|
||||
|
||||
# 手动模式按钮
|
||||
self.manual_btn = ttk.Button(self.btn_frame, text="按住说话")
|
||||
self.manual_btn.bind("<ButtonPress-1>", self._on_manual_button_press)
|
||||
self.manual_btn.bind("<ButtonRelease-1>", self._on_manual_button_release)
|
||||
self.manual_btn.pack(side=tk.LEFT, padx=10)
|
||||
|
||||
# 打断按钮
|
||||
self.abort_btn = ttk.Button(
|
||||
self.btn_frame, text="停止播放", command=self._on_abort_button_click
|
||||
)
|
||||
self.abort_btn.pack(side=tk.LEFT, padx=10)
|
||||
|
||||
# 对话模式标志
|
||||
self.auto_mode = False
|
||||
|
||||
# 回调函数
|
||||
self.button_press_callback = None
|
||||
self.button_release_callback = None
|
||||
self.status_update_callback = None
|
||||
self.text_update_callback = None
|
||||
self.emotion_update_callback = None
|
||||
self.mode_callback = None
|
||||
self.auto_callback = None
|
||||
self.abort_callback = None
|
||||
|
||||
# 更新队列
|
||||
self.update_queue = queue.Queue()
|
||||
|
||||
# 运行标志
|
||||
self._running = True
|
||||
|
||||
# 设置窗口关闭处理
|
||||
self.root.protocol("WM_DELETE_WINDOW", self.on_close)
|
||||
|
||||
# 启动更新处理
|
||||
self.root.after(100, self._process_updates)
|
||||
|
||||
def set_callbacks(
|
||||
self,
|
||||
press_callback: Optional[Callable] = None,
|
||||
release_callback: Optional[Callable] = None,
|
||||
status_callback: Optional[Callable] = None,
|
||||
text_callback: Optional[Callable] = None,
|
||||
emotion_callback: Optional[Callable] = None,
|
||||
mode_callback: Optional[Callable] = None,
|
||||
auto_callback: Optional[Callable] = None,
|
||||
abort_callback: Optional[Callable] = None,
|
||||
):
|
||||
"""设置回调函数"""
|
||||
self.button_press_callback = press_callback
|
||||
self.button_release_callback = release_callback
|
||||
self.status_update_callback = status_callback
|
||||
self.text_update_callback = text_callback
|
||||
self.emotion_update_callback = emotion_callback
|
||||
self.mode_callback = mode_callback
|
||||
self.auto_callback = auto_callback
|
||||
self.abort_callback = abort_callback
|
||||
|
||||
def _process_updates(self):
|
||||
"""处理更新队列"""
|
||||
try:
|
||||
while True:
|
||||
try:
|
||||
# 非阻塞方式获取更新
|
||||
update_func = self.update_queue.get_nowait()
|
||||
update_func()
|
||||
self.update_queue.task_done()
|
||||
except queue.Empty:
|
||||
break
|
||||
finally:
|
||||
if self._running:
|
||||
self.root.after(100, self._process_updates)
|
||||
|
||||
def _on_manual_button_press(self, event):
|
||||
"""手动模式按钮按下事件处理"""
|
||||
try:
|
||||
# 更新按钮文本为"松开以停止"
|
||||
self.manual_btn.config(text="松开以停止")
|
||||
|
||||
# 调用回调函数
|
||||
if self.button_press_callback:
|
||||
self.button_press_callback()
|
||||
except Exception as e:
|
||||
self.logger.error(f"按钮按下回调执行失败: {e}")
|
||||
|
||||
def _on_manual_button_release(self, event):
|
||||
"""手动模式按钮释放事件处理"""
|
||||
try:
|
||||
# 更新按钮文本为"按住说话"
|
||||
self.manual_btn.config(text="按住说话")
|
||||
|
||||
# 调用回调函数
|
||||
if self.button_release_callback:
|
||||
self.button_release_callback()
|
||||
except Exception as e:
|
||||
self.logger.error(f"按钮释放回调执行失败: {e}")
|
||||
|
||||
def _on_auto_button_click(self):
|
||||
"""自动模式按钮点击事件处理"""
|
||||
try:
|
||||
if self.auto_callback:
|
||||
self.auto_callback()
|
||||
except Exception as e:
|
||||
self.logger.error(f"自动模式按钮回调执行失败: {e}")
|
||||
|
||||
def _on_abort_button_click(self):
|
||||
"""打断按钮点击事件处理"""
|
||||
try:
|
||||
if self.abort_callback:
|
||||
self.abort_callback()
|
||||
except Exception as e:
|
||||
self.logger.error(f"打断按钮回调执行失败: {e}")
|
||||
|
||||
def _on_mode_button_click(self):
|
||||
"""对话模式切换按钮点击事件"""
|
||||
try:
|
||||
# 检查是否可以切换模式(通过回调函数询问应用程序当前状态)
|
||||
if self.mode_callback:
|
||||
# 如果回调函数返回False,表示当前不能切换模式
|
||||
if not self.mode_callback(not self.auto_mode):
|
||||
return
|
||||
|
||||
# 切换模式
|
||||
self.auto_mode = not self.auto_mode
|
||||
|
||||
# 更新按钮显示
|
||||
if self.auto_mode:
|
||||
# 切换到自动模式
|
||||
self.update_mode_button_status("自动对话")
|
||||
|
||||
# 隐藏手动按钮,显示自动按钮
|
||||
self.update_queue.put(lambda: self._switch_to_auto_mode())
|
||||
else:
|
||||
# 切换到手动模式
|
||||
self.update_mode_button_status("手动对话")
|
||||
|
||||
# 隐藏自动按钮,显示手动按钮
|
||||
self.update_queue.put(lambda: self._switch_to_manual_mode())
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"模式切换按钮回调执行失败: {e}")
|
||||
|
||||
def _switch_to_auto_mode(self):
|
||||
"""切换到自动模式的UI更新"""
|
||||
self.manual_btn.pack_forget() # 移除手动按钮
|
||||
self.auto_btn.pack(
|
||||
side=tk.LEFT, padx=10, before=self.abort_btn
|
||||
) # 显示自动按钮,放在打断按钮前面
|
||||
|
||||
def _switch_to_manual_mode(self):
|
||||
"""切换到手动模式的UI更新"""
|
||||
self.auto_btn.pack_forget() # 移除自动按钮
|
||||
self.manual_btn.pack(
|
||||
side=tk.LEFT, padx=10, before=self.abort_btn
|
||||
) # 显示手动按钮,放在打断按钮前面
|
||||
|
||||
def update_status(self, status: str):
|
||||
"""更新状态文本"""
|
||||
self.update_queue.put(lambda: self.status_label.config(text=f"状态: {status}"))
|
||||
|
||||
def update_text(self, text: str):
|
||||
"""更新TTS文本"""
|
||||
self.update_queue.put(lambda: self.tts_text_label.config(text=text))
|
||||
|
||||
def update_emotion(self, emotion: str):
|
||||
"""更新表情"""
|
||||
self.update_queue.put(lambda: self.emotion_label.config(text=emotion))
|
||||
|
||||
def start_update_threads(self):
|
||||
"""启动更新线程"""
|
||||
|
||||
def update_loop():
|
||||
while self._running:
|
||||
try:
|
||||
# 更新状态
|
||||
if self.status_update_callback:
|
||||
status = self.status_update_callback()
|
||||
if status:
|
||||
self.update_status(status)
|
||||
|
||||
# 更新文本
|
||||
if self.text_update_callback:
|
||||
text = self.text_update_callback()
|
||||
if text:
|
||||
self.update_text(text)
|
||||
|
||||
# 更新表情
|
||||
if self.emotion_update_callback:
|
||||
emotion = self.emotion_update_callback()
|
||||
if emotion:
|
||||
self.update_emotion(emotion)
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"更新失败: {e}")
|
||||
time.sleep(0.1)
|
||||
|
||||
threading.Thread(target=update_loop, daemon=True).start()
|
||||
|
||||
def on_close(self):
|
||||
"""关闭窗口处理"""
|
||||
self._running = False
|
||||
self.root.destroy()
|
||||
|
||||
def start(self):
|
||||
"""启动GUI"""
|
||||
try:
|
||||
# 启动更新线程
|
||||
self.start_update_threads()
|
||||
# 在主线程中运行主循环
|
||||
self.logger.info("开始启动GUI主循环")
|
||||
self.root.mainloop()
|
||||
except Exception as e:
|
||||
self.logger.error(f"GUI启动失败: {e}", exc_info=True)
|
||||
# 尝试回退到CLI模式
|
||||
print(f"GUI启动失败: {e},请尝试使用CLI模式")
|
||||
|
||||
def update_mode_button_status(self, text: str):
|
||||
"""更新模式按钮状态"""
|
||||
self.update_queue.put(lambda: self.mode_btn.config(text=text))
|
||||
|
||||
def update_button_status(self, text: str):
|
||||
"""更新按钮状态 - 保留此方法以满足抽象基类要求"""
|
||||
# 根据当前模式更新相应的按钮
|
||||
if self.auto_mode:
|
||||
self.update_queue.put(lambda: self.auto_btn.config(text=text))
|
||||
else:
|
||||
# 在手动模式下,不通过此方法更新按钮文本
|
||||
# 因为按钮文本由按下/释放事件直接控制
|
||||
pass
|
||||
@@ -0,0 +1,83 @@
|
||||
import json
|
||||
|
||||
from xiaozhi.services.protocols.typing import AbortReason, ListeningMode
|
||||
|
||||
|
||||
class Protocol:
|
||||
def __init__(self):
|
||||
self.session_id = ""
|
||||
self.on_incoming_json = None
|
||||
self.on_incoming_audio = None
|
||||
self.on_audio_channel_opened = None
|
||||
self.on_audio_channel_closed = None
|
||||
self.on_network_error = None
|
||||
|
||||
def on_incoming_json(self, callback):
|
||||
"""设置JSON消息接收回调函数"""
|
||||
self.on_incoming_json = callback
|
||||
|
||||
def on_incoming_audio(self, callback):
|
||||
"""设置音频数据接收回调函数"""
|
||||
self.on_incoming_audio = callback
|
||||
|
||||
def on_audio_channel_opened(self, callback):
|
||||
"""设置音频通道打开回调函数"""
|
||||
self.on_audio_channel_opened = callback
|
||||
|
||||
def on_audio_channel_closed(self, callback):
|
||||
"""设置音频通道关闭回调函数"""
|
||||
self.on_audio_channel_closed = callback
|
||||
|
||||
def on_network_error(self, callback):
|
||||
"""设置网络错误回调函数"""
|
||||
self.on_network_error = callback
|
||||
|
||||
async def send_text(self, message):
|
||||
"""发送文本消息的抽象方法,需要在子类中实现"""
|
||||
raise NotImplementedError("send_text方法必须由子类实现")
|
||||
|
||||
async def send_abort_speaking(self, reason):
|
||||
"""发送中止语音的消息"""
|
||||
message = {"session_id": self.session_id, "type": "abort"}
|
||||
if reason == AbortReason.WAKE_WORD_DETECTED:
|
||||
message["reason"] = "wake_word_detected"
|
||||
await self.send_text(json.dumps(message))
|
||||
|
||||
|
||||
async def send_start_listening(self, mode):
|
||||
"""发送开始监听的消息"""
|
||||
mode_map = {
|
||||
ListeningMode.ALWAYS_ON: "realtime",
|
||||
ListeningMode.AUTO_STOP: "auto",
|
||||
ListeningMode.MANUAL: "manual",
|
||||
}
|
||||
message = {
|
||||
"session_id": self.session_id,
|
||||
"type": "listen",
|
||||
"state": "start",
|
||||
"mode": mode_map[mode],
|
||||
}
|
||||
await self.send_text(json.dumps(message))
|
||||
|
||||
async def send_stop_listening(self):
|
||||
"""发送停止监听的消息"""
|
||||
message = {"session_id": self.session_id, "type": "listen", "state": "stop"}
|
||||
await self.send_text(json.dumps(message))
|
||||
|
||||
async def send_iot_descriptors(self, descriptors):
|
||||
"""发送物联网设备描述信息"""
|
||||
message = {
|
||||
"session_id": self.session_id,
|
||||
"type": "iot",
|
||||
"descriptors": json.loads(descriptors),
|
||||
}
|
||||
await self.send_text(json.dumps(message))
|
||||
|
||||
async def send_iot_states(self, states):
|
||||
"""发送物联网设备状态信息"""
|
||||
message = {
|
||||
"session_id": self.session_id,
|
||||
"type": "iot",
|
||||
"states": json.loads(states),
|
||||
}
|
||||
await self.send_text(json.dumps(message))
|
||||
@@ -0,0 +1,30 @@
|
||||
class ListeningMode:
|
||||
"""监听模式"""
|
||||
ALWAYS_ON = "always_on"
|
||||
AUTO_STOP = "auto_stop"
|
||||
MANUAL = "manual"
|
||||
|
||||
class AbortReason:
|
||||
"""中止原因"""
|
||||
NONE = "none"
|
||||
WAKE_WORD_DETECTED = "wake_word_detected"
|
||||
|
||||
class DeviceState:
|
||||
"""设备状态"""
|
||||
IDLE = "idle"
|
||||
CONNECTING = "connecting"
|
||||
LISTENING = "listening"
|
||||
SPEAKING = "speaking"
|
||||
|
||||
class EventType:
|
||||
"""事件类型"""
|
||||
SCHEDULE_EVENT = "schedule_event"
|
||||
AUDIO_INPUT_READY_EVENT = "audio_input_ready_event"
|
||||
AUDIO_OUTPUT_READY_EVENT = "audio_output_ready_event"
|
||||
|
||||
class AudioConfig:
|
||||
"""音频配置"""
|
||||
SAMPLE_RATE = 24000
|
||||
CHANNELS = 1
|
||||
FRAME_DURATION = 60 # ms
|
||||
FRAME_SIZE = int(SAMPLE_RATE * (FRAME_DURATION / 1000))
|
||||
@@ -0,0 +1,214 @@
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import websockets
|
||||
|
||||
|
||||
from xiaozhi.services.protocols.protocol import Protocol
|
||||
from xiaozhi.utils.config_manager import ConfigManager
|
||||
|
||||
|
||||
logger = logging.getLogger("WebsocketProtocol")
|
||||
|
||||
|
||||
class WebsocketProtocol(Protocol):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
# 获取配置管理器实例
|
||||
self.config = ConfigManager.instance()
|
||||
self.websocket = None
|
||||
self.server_sample_rate = 16000
|
||||
self.connected = False
|
||||
self.hello_received = None # 初始化时先设为 None
|
||||
self.WEBSOCKET_URL = self.config.get_config("NETWORK.WEBSOCKET_URL")
|
||||
self.WEBSOCKET_ACCESS_TOKEN = self.config.get_config(
|
||||
"NETWORK.WEBSOCKET_ACCESS_TOKEN"
|
||||
)
|
||||
self.CLIENT_ID = self.config.get_client_id()
|
||||
self.DEVICE_ID = self.config.get_device_id()
|
||||
|
||||
async def connect(self) -> bool:
|
||||
"""连接到WebSocket服务器"""
|
||||
try:
|
||||
# 在连接时创建 Event,确保在正确的事件循环中
|
||||
self.hello_received = asyncio.Event()
|
||||
|
||||
# 配置连接
|
||||
headers = {
|
||||
"Authorization": f"Bearer {self.WEBSOCKET_ACCESS_TOKEN}",
|
||||
"Protocol-Version": "1",
|
||||
"Device-Id": self.DEVICE_ID, # 获取设备MAC地址
|
||||
"Client-Id": self.CLIENT_ID,
|
||||
}
|
||||
|
||||
# 建立WebSocket连接 (兼容不同Python版本的写法)
|
||||
try:
|
||||
# 新的写法 (在Python 3.11+版本中)
|
||||
self.websocket = await websockets.connect(
|
||||
uri=self.WEBSOCKET_URL, additional_headers=headers
|
||||
)
|
||||
except TypeError:
|
||||
# 旧的写法 (在较早的Python版本中)
|
||||
self.websocket = await websockets.connect(
|
||||
self.WEBSOCKET_URL, extra_headers=headers
|
||||
)
|
||||
|
||||
# 启动消息处理循环
|
||||
asyncio.create_task(self._message_handler())
|
||||
|
||||
# 发送客户端hello消息
|
||||
hello_message = {
|
||||
"type": "hello",
|
||||
"version": 1,
|
||||
"transport": "websocket",
|
||||
"audio_params": {
|
||||
"format": "opus",
|
||||
"sample_rate": 16000,
|
||||
"channels": 1,
|
||||
"frame_duration": 60,
|
||||
},
|
||||
}
|
||||
await self.send_text(json.dumps(hello_message))
|
||||
|
||||
# 等待服务器hello响应
|
||||
try:
|
||||
await asyncio.wait_for(self.hello_received.wait(), timeout=10.0)
|
||||
self.connected = True
|
||||
logger.info("已连接到WebSocket服务器")
|
||||
return True
|
||||
except asyncio.TimeoutError:
|
||||
logger.error("等待服务器hello响应超时")
|
||||
if self.on_network_error:
|
||||
self.on_network_error("等待响应超时")
|
||||
return False
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"WebSocket连接失败: {e}")
|
||||
if self.on_network_error:
|
||||
self.on_network_error(f"无法连接服务: {str(e)}")
|
||||
return False
|
||||
|
||||
async def _message_handler(self):
|
||||
"""处理接收到的WebSocket消息"""
|
||||
try:
|
||||
async for message in self.websocket:
|
||||
if isinstance(message, str):
|
||||
try:
|
||||
data = json.loads(message)
|
||||
msg_type = data.get("type")
|
||||
if msg_type == "hello":
|
||||
# 处理服务器 hello 消息
|
||||
await self._handle_server_hello(data)
|
||||
else:
|
||||
if self.on_incoming_json:
|
||||
self.on_incoming_json(data)
|
||||
except json.JSONDecodeError as e:
|
||||
logger.error(f"无效的JSON消息: {message}, 错误: {e}")
|
||||
elif self.on_incoming_audio: # 使用 elif 更清晰
|
||||
self.on_incoming_audio(message)
|
||||
|
||||
except websockets.ConnectionClosed:
|
||||
logger.info("WebSocket连接已关闭")
|
||||
self.connected = False
|
||||
if self.on_audio_channel_closed:
|
||||
# 使用 schedule 确保回调在主线程中执行
|
||||
await self.on_audio_channel_closed()
|
||||
except Exception as e:
|
||||
logger.error(f"消息处理错误: {e}")
|
||||
self.connected = False
|
||||
if self.on_network_error:
|
||||
# 使用 schedule 确保错误处理在主线程中执行
|
||||
self.on_network_error(f"连接错误: {str(e)}")
|
||||
|
||||
async def send_audio(self, data: bytes):
|
||||
"""发送音频数据"""
|
||||
if not self.is_audio_channel_opened(): # 使用已有的 is_connected 方法
|
||||
return
|
||||
|
||||
try:
|
||||
await self.websocket.send(data)
|
||||
except Exception as e:
|
||||
logger.error(f"发送音频数据失败: {e}")
|
||||
if self.on_network_error:
|
||||
self.on_network_error(f"发送音频失败: {str(e)}")
|
||||
|
||||
async def send_text(self, message: str):
|
||||
"""发送文本消息"""
|
||||
if self.websocket:
|
||||
try:
|
||||
await self.websocket.send(message)
|
||||
except Exception as e:
|
||||
await self.close_audio_channel()
|
||||
if self.on_network_error:
|
||||
self.on_network_error(f"发送消息失败: {str(e)}")
|
||||
|
||||
def is_audio_channel_opened(self) -> bool:
|
||||
"""检查音频通道是否打开"""
|
||||
return self.websocket is not None and self.connected
|
||||
|
||||
async def open_audio_channel(self) -> bool:
|
||||
"""建立 WebSocket 连接
|
||||
|
||||
如果尚未连接,则创建新的 WebSocket 连接
|
||||
Returns:
|
||||
bool: 连接是否成功
|
||||
"""
|
||||
if not self.connected:
|
||||
return await self.connect()
|
||||
return True
|
||||
|
||||
async def _handle_server_hello(self, data: dict):
|
||||
"""处理服务器的 hello 消息
|
||||
|
||||
解析服务器返回的 hello 消息,设置相关参数并通知音频通道已打开
|
||||
|
||||
Args:
|
||||
data: 服务器返回的 hello 消息数据
|
||||
"""
|
||||
try:
|
||||
# 验证传输方式
|
||||
transport = data.get("transport")
|
||||
if not transport or transport != "websocket":
|
||||
logger.error(f"不支持的传输方式: {transport}")
|
||||
return
|
||||
|
||||
# 获取音频参数
|
||||
audio_params = data.get("audio_params")
|
||||
if audio_params:
|
||||
# 获取服务器的采样率
|
||||
sample_rate = audio_params.get("sample_rate")
|
||||
if sample_rate:
|
||||
self.server_sample_rate = sample_rate
|
||||
# 如果服务器采样率与本地不同,记录警告
|
||||
if sample_rate != self.server_sample_rate:
|
||||
logger.warning(
|
||||
f"服务器的音频采样率 {sample_rate} "
|
||||
f"与设备输出的采样率 {self.server_sample_rate} 不一致,"
|
||||
"重采样后可能会失真"
|
||||
)
|
||||
|
||||
# 设置 hello 接收事件
|
||||
self.hello_received.set()
|
||||
|
||||
# 通知音频通道已打开
|
||||
if self.on_audio_channel_opened:
|
||||
await self.on_audio_channel_opened()
|
||||
|
||||
logger.info("成功处理服务器 hello 消息")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"处理服务器 hello 消息时出错: {e}")
|
||||
if self.on_network_error:
|
||||
self.on_network_error(f"处理服务器响应失败: {str(e)}")
|
||||
|
||||
async def close_audio_channel(self):
|
||||
"""关闭音频通道"""
|
||||
if self.websocket:
|
||||
try:
|
||||
await self.websocket.close()
|
||||
self.websocket = None
|
||||
self.connected = False
|
||||
if self.on_audio_channel_closed:
|
||||
await self.on_audio_channel_closed()
|
||||
except Exception as e:
|
||||
logger.error(f"关闭WebSocket连接失败: {e}")
|
||||
@@ -0,0 +1,256 @@
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Dict, Any, Optional
|
||||
import threading
|
||||
import requests
|
||||
import socket
|
||||
import uuid
|
||||
|
||||
logger = logging.getLogger("ConfigManager")
|
||||
|
||||
|
||||
class ConfigManager:
|
||||
"""配置管理器 - 单例模式"""
|
||||
|
||||
_instance = None
|
||||
_lock = threading.Lock()
|
||||
CONFIG_FILE = Path(os.getcwd()) / "xiaozhi.json"
|
||||
|
||||
# 默认配置
|
||||
DEFAULT_CONFIG = {
|
||||
"CLIENT_ID": None,
|
||||
"DEVICE_ID": None,
|
||||
"NETWORK": {
|
||||
"OTA_VERSION_URL": "https://api.tenclass.net/xiaozhi/ota/",
|
||||
"WEBSOCKET_URL": "wss://api.tenclass.net/xiaozhi/v1/",
|
||||
"WEBSOCKET_ACCESS_TOKEN": "test-token",
|
||||
},
|
||||
"MQTT_INFO": None,
|
||||
}
|
||||
|
||||
def __new__(cls):
|
||||
"""确保单例模式"""
|
||||
if cls._instance is None:
|
||||
cls._instance = super().__new__(cls)
|
||||
return cls._instance
|
||||
|
||||
def __init__(self):
|
||||
"""初始化配置管理器"""
|
||||
self.logger = logger
|
||||
if hasattr(self, "_initialized"):
|
||||
return
|
||||
self._initialized = True
|
||||
|
||||
# 加载配置
|
||||
self._config = self._load_config()
|
||||
self._initialize_client_id()
|
||||
self._initialize_device_id()
|
||||
self._initialize_mqtt_info()
|
||||
|
||||
def _load_config(self) -> Dict[str, Any]:
|
||||
"""加载配置文件,如果不存在则创建"""
|
||||
try:
|
||||
if self.CONFIG_FILE.exists():
|
||||
config = json.loads(self.CONFIG_FILE.read_text(encoding="utf-8"))
|
||||
return self._merge_configs(self.DEFAULT_CONFIG, config)
|
||||
else:
|
||||
self._save_config(self.DEFAULT_CONFIG)
|
||||
return self.DEFAULT_CONFIG.copy()
|
||||
except Exception as e:
|
||||
logger.error(f"Error loading config: {e}")
|
||||
return self.DEFAULT_CONFIG.copy()
|
||||
|
||||
def _save_config(self, config: dict) -> bool:
|
||||
"""保存配置到文件"""
|
||||
try:
|
||||
self.CONFIG_FILE.write_text(
|
||||
json.dumps(config, indent=2, ensure_ascii=False), encoding="utf-8"
|
||||
)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"Error saving config: {e}")
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def _merge_configs(default: dict, custom: dict) -> dict:
|
||||
"""递归合并配置字典"""
|
||||
result = default.copy()
|
||||
for key, value in custom.items():
|
||||
if (
|
||||
key in result
|
||||
and isinstance(result[key], dict)
|
||||
and isinstance(value, dict)
|
||||
):
|
||||
result[key] = ConfigManager._merge_configs(result[key], value)
|
||||
else:
|
||||
result[key] = value
|
||||
return result
|
||||
|
||||
def get_client_id(self) -> str:
|
||||
"""获取客户端ID"""
|
||||
return self._config["CLIENT_ID"]
|
||||
|
||||
def get_device_id(self) -> Optional[str]:
|
||||
"""获取设备ID"""
|
||||
return self._config.get("DEVICE_ID")
|
||||
|
||||
def get_network_config(self) -> dict:
|
||||
"""获取网络配置"""
|
||||
return self._config["NETWORK"]
|
||||
|
||||
def get_config(self, path: str, default: Any = None) -> Any:
|
||||
"""
|
||||
通过路径获取配置值
|
||||
"""
|
||||
try:
|
||||
value = self._config
|
||||
for key in path.split("."):
|
||||
value = value[key]
|
||||
return value
|
||||
except (KeyError, TypeError):
|
||||
return default
|
||||
|
||||
def update_config(self, path: str, value: Any) -> bool:
|
||||
"""
|
||||
更新特定配置项
|
||||
"""
|
||||
try:
|
||||
current = self._config
|
||||
*parts, last = path.split(".")
|
||||
for part in parts:
|
||||
current = current.setdefault(part, {})
|
||||
current[last] = value
|
||||
return self._save_config(self._config)
|
||||
except Exception as e:
|
||||
logger.error(f"Error updating config {path}: {e}")
|
||||
return False
|
||||
|
||||
@classmethod
|
||||
def instance(cls):
|
||||
"""获取配置管理器实例(线程安全)"""
|
||||
with cls._lock:
|
||||
if cls._instance is None:
|
||||
cls._instance = cls()
|
||||
return cls._instance
|
||||
|
||||
def get_mac_address(self):
|
||||
mac = uuid.UUID(int=uuid.getnode()).hex[-12:]
|
||||
return ":".join([mac[i : i + 2] for i in range(0, 12, 2)])
|
||||
|
||||
def generate_uuid(self) -> str:
|
||||
return str(uuid.uuid4())
|
||||
|
||||
def get_local_ip(self):
|
||||
try:
|
||||
s = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
|
||||
s.connect(("8.8.8.8", 80))
|
||||
ip = s.getsockname()[0]
|
||||
s.close()
|
||||
return ip
|
||||
except Exception:
|
||||
return "127.0.0.1"
|
||||
|
||||
def _initialize_client_id(self):
|
||||
"""确保存在客户端ID"""
|
||||
if not self._config["CLIENT_ID"]:
|
||||
client_id = self.generate_uuid()
|
||||
success = self.update_config("CLIENT_ID", client_id)
|
||||
if success:
|
||||
logger.info(f"Generated new CLIENT_ID: {client_id}")
|
||||
else:
|
||||
logger.error("Failed to save new CLIENT_ID")
|
||||
|
||||
def _initialize_device_id(self):
|
||||
"""确保存在设备ID"""
|
||||
if not self._config["DEVICE_ID"]:
|
||||
try:
|
||||
device_hash = self.get_mac_address()
|
||||
success = self.update_config("DEVICE_ID", device_hash)
|
||||
if success:
|
||||
logger.info(f"Generated new DEVICE_ID: {device_hash}")
|
||||
else:
|
||||
logger.error("Failed to save new DEVICE_ID")
|
||||
except Exception as e:
|
||||
logger.error(f"Error generating DEVICE_ID: {e}")
|
||||
|
||||
def _initialize_mqtt_info(self):
|
||||
try:
|
||||
mqtt_info = self._get_ota_version()
|
||||
if mqtt_info:
|
||||
self.update_config("MQTT_INFO", mqtt_info)
|
||||
self.logger.info("MQTT信息已成功更新")
|
||||
return mqtt_info
|
||||
else:
|
||||
self.logger.warning("获取MQTT信息失败,使用已保存的配置")
|
||||
return self.get_config("MQTT_INFO")
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"初始化MQTT信息失败: {e}")
|
||||
return self.get_config("MQTT_INFO")
|
||||
|
||||
def _get_ota_version(self):
|
||||
"""获取OTA服务器的MQTT信息"""
|
||||
MAC_ADDR = self.get_device_id()
|
||||
OTA_VERSION_URL = self.get_config("NETWORK.OTA_VERSION_URL")
|
||||
headers = {"Device-Id": MAC_ADDR, "Content-Type": "application/json"}
|
||||
|
||||
# 构建设备信息payload
|
||||
payload = {
|
||||
"flash_size": 16777216, # 闪存大小 (16MB)
|
||||
"minimum_free_heap_size": 8318916, # 最小可用堆内存
|
||||
"mac_address": MAC_ADDR, # 设备MAC地址
|
||||
"chip_model_name": "esp32s3", # 芯片型号
|
||||
"chip_info": {"model": 9, "cores": 2, "revision": 2, "features": 18},
|
||||
"application": {
|
||||
"name": "xiaozhi",
|
||||
"version": "1.1.2",
|
||||
"idf_version": "v5.3.2-dirty",
|
||||
},
|
||||
"partition_table": [],
|
||||
"ota": {"label": "factory"},
|
||||
"board": {
|
||||
"type": "bread-compact-wifi",
|
||||
"ip": self.get_local_ip(),
|
||||
"mac": MAC_ADDR,
|
||||
},
|
||||
}
|
||||
|
||||
try:
|
||||
# 发送请求到OTA服务器
|
||||
response = requests.post(
|
||||
OTA_VERSION_URL,
|
||||
headers=headers,
|
||||
json=payload,
|
||||
timeout=10,
|
||||
)
|
||||
|
||||
# 检查HTTP状态码
|
||||
if response.status_code != 200:
|
||||
self.logger.error(f"OTA服务器错误: HTTP {response.status_code}")
|
||||
raise ValueError(f"OTA服务器返回错误状态码: {response.status_code}")
|
||||
|
||||
# 解析JSON数据
|
||||
response_data = response.json()
|
||||
|
||||
# 调试信息:打印完整的OTA响应
|
||||
self.logger.debug(
|
||||
f"OTA服务器返回数据: {json.dumps(response_data, indent=4, ensure_ascii=False)}"
|
||||
)
|
||||
|
||||
# 确保"mqtt"信息存在
|
||||
if "mqtt" in response_data:
|
||||
self.logger.info(f"MQTT服务器信息已更新")
|
||||
return response_data["mqtt"]
|
||||
else:
|
||||
self.logger.error("OTA服务器返回的数据无效: MQTT信息缺失")
|
||||
raise ValueError("OTA服务器返回的数据无效,请检查服务器状态或MAC地址!")
|
||||
|
||||
except requests.Timeout:
|
||||
self.logger.error("OTA请求超时,请检查网络或服务器状态")
|
||||
raise ValueError("OTA请求超时!请稍后重试。")
|
||||
|
||||
except requests.RequestException as e:
|
||||
self.logger.error(f"OTA请求失败: {e}")
|
||||
raise ValueError("无法连接到OTA服务器,请检查网络连接!")
|
||||
@@ -0,0 +1,30 @@
|
||||
import logging
|
||||
|
||||
|
||||
def setup_logging():
|
||||
"""配置日志系统"""
|
||||
|
||||
# 创建根日志记录器
|
||||
root_logger = logging.getLogger()
|
||||
root_logger.setLevel(logging.INFO) # 设置根日志级别
|
||||
|
||||
# 清除已有的处理器(避免重复添加)
|
||||
if root_logger.handlers:
|
||||
root_logger.handlers.clear()
|
||||
|
||||
# 创建控制台处理器
|
||||
console_handler = logging.StreamHandler()
|
||||
console_handler.setLevel(logging.INFO)
|
||||
|
||||
# 创建格式化器
|
||||
formatter = logging.Formatter(
|
||||
"%(asctime)s - %(name)s - %(levelname)s - %(message)s"
|
||||
)
|
||||
console_handler.setFormatter(formatter)
|
||||
|
||||
# 添加处理器到根日志记录器
|
||||
root_logger.addHandler(console_handler)
|
||||
|
||||
# 设置特定模块的日志级别
|
||||
logging.getLogger("XiaoZhi").setLevel(logging.INFO)
|
||||
logging.getLogger("WebsocketProtocol").setLevel(logging.INFO)
|
||||
@@ -0,0 +1,47 @@
|
||||
import argparse
|
||||
import asyncio
|
||||
|
||||
import numpy as np
|
||||
import open_xiaoai_server
|
||||
|
||||
from xiaozhi.services.audio.stream import GlobalStream
|
||||
|
||||
|
||||
class XiaoAi:
|
||||
mode = "xiaoai"
|
||||
loop = asyncio.new_event_loop()
|
||||
|
||||
@classmethod
|
||||
def setup_mode(cls):
|
||||
parser = argparse.ArgumentParser(
|
||||
description="小爱音箱接入小智 AI 演示 | by: https://del.wang"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--mode",
|
||||
type=str,
|
||||
choices=["xiaoai", "xiaozhi"],
|
||||
default="xiaoai",
|
||||
help="运行模式:【xiaoai】使用小爱音箱的输入输出音频(默认)、【xiaozhi】使用本地电脑的输入输出音频",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
if args.mode == "xiaozhi":
|
||||
cls.mode = "xiaozhi"
|
||||
|
||||
@classmethod
|
||||
def on_input_data(cls, data: bytes):
|
||||
audio_array = np.frombuffer(data, dtype=np.uint16)
|
||||
GlobalStream().add_input_data(audio_array.tobytes())
|
||||
|
||||
@classmethod
|
||||
def on_output_data(cls, data: bytes):
|
||||
async def on_output_data_async(data: bytes):
|
||||
return await open_xiaoai_server.on_output_data(data)
|
||||
|
||||
future = on_output_data_async(data)
|
||||
cls.loop.run_until_complete(future)
|
||||
|
||||
@classmethod
|
||||
async def init_xiaoai(cls):
|
||||
GlobalStream().on_output_data = cls.on_output_data
|
||||
open_xiaoai_server.register_fn("on_input_data", cls.on_input_data)
|
||||
await open_xiaoai_server.start_server()
|
||||
@@ -0,0 +1,808 @@
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
|
||||
from xiaozhi.services.display import gui_display
|
||||
from xiaozhi.services.protocols.typing import (
|
||||
AbortReason,
|
||||
AudioConfig,
|
||||
DeviceState,
|
||||
EventType,
|
||||
ListeningMode,
|
||||
)
|
||||
from xiaozhi.services.protocols.websocket_protocol import WebsocketProtocol
|
||||
from xiaozhi.utils.config_manager import ConfigManager
|
||||
from xiaozhi.xiaoai import XiaoAi
|
||||
|
||||
# 配置日志
|
||||
logger = logging.getLogger("XiaoZhi")
|
||||
|
||||
|
||||
class XiaoZhi:
|
||||
"""智能音箱应用程序主类"""
|
||||
|
||||
_instance = None
|
||||
|
||||
@classmethod
|
||||
def instance(cls):
|
||||
"""获取单例实例"""
|
||||
if cls._instance is None:
|
||||
cls._instance = XiaoZhi()
|
||||
return cls._instance
|
||||
|
||||
def __init__(self):
|
||||
"""初始化应用程序"""
|
||||
# 确保单例模式
|
||||
if XiaoZhi._instance is not None:
|
||||
raise Exception("XiaoZhi是单例类,请使用instance()获取实例")
|
||||
XiaoZhi._instance = self
|
||||
|
||||
# 获取配置管理器实例
|
||||
self.config = ConfigManager.instance()
|
||||
|
||||
# 状态变量
|
||||
self.device_state = DeviceState.IDLE
|
||||
self.voice_detected = False
|
||||
self.keep_listening = False
|
||||
self.aborted = False
|
||||
self.current_text = ""
|
||||
self.current_emotion = "neutral"
|
||||
|
||||
# 音频处理相关
|
||||
self.audio_codec = None
|
||||
|
||||
# 事件循环和线程
|
||||
self.loop = asyncio.new_event_loop()
|
||||
self.loop_thread = None
|
||||
self.running = False
|
||||
|
||||
# 任务队列和锁
|
||||
self.main_tasks = []
|
||||
self.mutex = threading.Lock()
|
||||
|
||||
# 协议实例
|
||||
self.protocol = None
|
||||
|
||||
# 回调函数
|
||||
self.on_state_changed_callbacks = []
|
||||
|
||||
# 初始化事件对象
|
||||
self.events = {
|
||||
EventType.SCHEDULE_EVENT: threading.Event(),
|
||||
EventType.AUDIO_INPUT_READY_EVENT: threading.Event(),
|
||||
EventType.AUDIO_OUTPUT_READY_EVENT: threading.Event(),
|
||||
}
|
||||
|
||||
# 创建显示界面
|
||||
self.display = None
|
||||
|
||||
def run(self):
|
||||
self.protocol = WebsocketProtocol()
|
||||
|
||||
# 创建并启动事件循环线程
|
||||
self.loop_thread = threading.Thread(target=self._run_event_loop)
|
||||
self.loop_thread.daemon = True
|
||||
self.loop_thread.start()
|
||||
|
||||
# 等待事件循环准备就绪
|
||||
time.sleep(0.1)
|
||||
|
||||
# 初始化应用程序(移除自动连接)
|
||||
asyncio.run_coroutine_threadsafe(XiaoAi.init_xiaoai(), self.loop)
|
||||
asyncio.run_coroutine_threadsafe(self._initialize_without_connect(), self.loop)
|
||||
|
||||
# 启动主循环线程
|
||||
main_loop_thread = threading.Thread(target=self._main_loop)
|
||||
main_loop_thread.daemon = True
|
||||
main_loop_thread.start()
|
||||
|
||||
# 启动 GUI
|
||||
self._initialize_display()
|
||||
self.display.start()
|
||||
|
||||
def _run_event_loop(self):
|
||||
"""运行事件循环的线程函数"""
|
||||
asyncio.set_event_loop(self.loop)
|
||||
self.loop.run_forever()
|
||||
|
||||
async def _initialize_without_connect(self):
|
||||
"""初始化应用程序组件(不建立连接)"""
|
||||
logger.info("正在初始化应用程序...")
|
||||
|
||||
# 设置设备状态为待命
|
||||
self.set_device_state(DeviceState.IDLE)
|
||||
|
||||
# 初始化音频编解码器
|
||||
self._initialize_audio()
|
||||
|
||||
# 设置协议回调
|
||||
self.protocol.on_network_error = self._on_network_error
|
||||
self.protocol.on_incoming_audio = self._on_incoming_audio
|
||||
self.protocol.on_incoming_json = self._on_incoming_json
|
||||
self.protocol.on_audio_channel_opened = self._on_audio_channel_opened
|
||||
self.protocol.on_audio_channel_closed = self._on_audio_channel_closed
|
||||
|
||||
logger.info("应用程序初始化完成")
|
||||
|
||||
def _initialize_audio(self):
|
||||
"""初始化音频设备和编解码器"""
|
||||
try:
|
||||
from xiaozhi.services.audio.codec import AudioCodec
|
||||
|
||||
self.audio_codec = AudioCodec()
|
||||
logger.info("音频编解码器初始化成功")
|
||||
except Exception as e:
|
||||
logger.error(f"初始化音频设备失败: {e}")
|
||||
self.alert("错误", f"初始化音频设备失败: {e}")
|
||||
|
||||
def _initialize_display(self):
|
||||
"""初始化显示界面"""
|
||||
self.display = gui_display.GuiDisplay()
|
||||
|
||||
# 设置回调函数
|
||||
self.display.set_callbacks(
|
||||
press_callback=self.start_listening,
|
||||
release_callback=self.stop_listening,
|
||||
status_callback=self._get_status_text,
|
||||
text_callback=self._get_current_text,
|
||||
emotion_callback=self._get_current_emotion,
|
||||
mode_callback=self._on_mode_changed,
|
||||
auto_callback=self.toggle_chat_state,
|
||||
abort_callback=lambda: self.abort_speaking(AbortReason.WAKE_WORD_DETECTED),
|
||||
)
|
||||
|
||||
def _main_loop(self):
|
||||
"""应用程序主循环"""
|
||||
logger.info("主循环已启动")
|
||||
self.running = True
|
||||
|
||||
while self.running:
|
||||
# 等待事件
|
||||
for event_type, event in self.events.items():
|
||||
if event.is_set():
|
||||
event.clear()
|
||||
|
||||
if event_type == EventType.AUDIO_INPUT_READY_EVENT:
|
||||
self._handle_input_audio()
|
||||
elif event_type == EventType.AUDIO_OUTPUT_READY_EVENT:
|
||||
self._handle_output_audio()
|
||||
elif event_type == EventType.SCHEDULE_EVENT:
|
||||
self._process_scheduled_tasks()
|
||||
|
||||
# 短暂休眠以避免CPU占用过高
|
||||
time.sleep(0.01)
|
||||
|
||||
def _process_scheduled_tasks(self):
|
||||
"""处理调度任务"""
|
||||
with self.mutex:
|
||||
tasks = self.main_tasks.copy()
|
||||
self.main_tasks.clear()
|
||||
|
||||
for task in tasks:
|
||||
try:
|
||||
task()
|
||||
except Exception as e:
|
||||
logger.error(f"执行调度任务时出错: {e}")
|
||||
|
||||
def schedule(self, callback):
|
||||
"""调度任务到主循环"""
|
||||
with self.mutex:
|
||||
# 如果是中止语音的任务,检查是否已经存在相同类型的任务
|
||||
if "abort_speaking" in str(callback):
|
||||
# 如果已经有中止任务在队列中,就不再添加
|
||||
if any("abort_speaking" in str(task) for task in self.main_tasks):
|
||||
return
|
||||
self.main_tasks.append(callback)
|
||||
self.events[EventType.SCHEDULE_EVENT].set()
|
||||
|
||||
def _handle_input_audio(self):
|
||||
"""处理音频输入"""
|
||||
if self.device_state != DeviceState.LISTENING:
|
||||
return
|
||||
|
||||
encoded_data = self.audio_codec.read_audio()
|
||||
if encoded_data and self.protocol and self.protocol.is_audio_channel_opened():
|
||||
asyncio.run_coroutine_threadsafe(
|
||||
self.protocol.send_audio(encoded_data), self.loop
|
||||
)
|
||||
|
||||
def _handle_output_audio(self):
|
||||
"""处理音频输出"""
|
||||
if self.device_state != DeviceState.SPEAKING:
|
||||
return
|
||||
|
||||
self.audio_codec.play_audio()
|
||||
|
||||
def _on_network_error(self, message):
|
||||
"""网络错误回调"""
|
||||
self.keep_listening = False
|
||||
self.set_device_state(DeviceState.IDLE)
|
||||
if self.device_state != DeviceState.CONNECTING:
|
||||
logger.info("检测到连接断开")
|
||||
self.set_device_state(DeviceState.IDLE)
|
||||
|
||||
# 关闭现有连接
|
||||
if self.protocol:
|
||||
asyncio.run_coroutine_threadsafe(
|
||||
self.protocol.close_audio_channel(), self.loop
|
||||
)
|
||||
|
||||
def _attempt_reconnect(self):
|
||||
"""尝试重新连接服务器"""
|
||||
if self.device_state != DeviceState.CONNECTING:
|
||||
logger.info("检测到连接断开,尝试重新连接...")
|
||||
self.set_device_state(DeviceState.CONNECTING)
|
||||
|
||||
# 关闭现有连接
|
||||
if self.protocol:
|
||||
asyncio.run_coroutine_threadsafe(
|
||||
self.protocol.close_audio_channel(), self.loop
|
||||
)
|
||||
|
||||
# 延迟一秒后尝试重新连接
|
||||
def delayed_reconnect():
|
||||
time.sleep(1)
|
||||
asyncio.run_coroutine_threadsafe(self._reconnect(), self.loop)
|
||||
|
||||
threading.Thread(target=delayed_reconnect, daemon=True).start()
|
||||
|
||||
async def _reconnect(self):
|
||||
"""重新连接到服务器"""
|
||||
|
||||
# 设置协议回调
|
||||
self.protocol.on_network_error = self._on_network_error
|
||||
self.protocol.on_incoming_audio = self._on_incoming_audio
|
||||
self.protocol.on_incoming_json = self._on_incoming_json
|
||||
self.protocol.on_audio_channel_opened = self._on_audio_channel_opened
|
||||
self.protocol.on_audio_channel_closed = self._on_audio_channel_closed
|
||||
|
||||
# 连接到服务器
|
||||
retry_count = 0
|
||||
max_retries = 3
|
||||
|
||||
while retry_count < max_retries:
|
||||
logger.info(f"尝试重新连接 (尝试 {retry_count + 1}/{max_retries})...")
|
||||
if await self.protocol.connect():
|
||||
logger.info("重新连接成功")
|
||||
self.set_device_state(DeviceState.IDLE)
|
||||
return True
|
||||
|
||||
retry_count += 1
|
||||
await asyncio.sleep(2) # 等待2秒后重试
|
||||
|
||||
logger.error(f"重新连接失败,已尝试 {max_retries} 次")
|
||||
self.schedule(lambda: self.alert("连接错误", "无法重新连接到服务器"))
|
||||
self.set_device_state(DeviceState.IDLE)
|
||||
return False
|
||||
|
||||
def _on_incoming_audio(self, data):
|
||||
"""接收音频数据回调"""
|
||||
if self.device_state == DeviceState.SPEAKING:
|
||||
self.audio_codec.write_audio(data)
|
||||
self.events[EventType.AUDIO_OUTPUT_READY_EVENT].set()
|
||||
|
||||
def _on_incoming_json(self, json_data):
|
||||
"""接收JSON数据回调"""
|
||||
try:
|
||||
if not json_data:
|
||||
return
|
||||
|
||||
# 解析JSON数据
|
||||
if isinstance(json_data, str):
|
||||
data = json.loads(json_data)
|
||||
else:
|
||||
data = json_data
|
||||
|
||||
# 处理不同类型的消息
|
||||
msg_type = data.get("type", "")
|
||||
if msg_type == "tts":
|
||||
self._handle_tts_message(data)
|
||||
elif msg_type == "stt":
|
||||
self._handle_stt_message(data)
|
||||
elif msg_type == "llm":
|
||||
self._handle_llm_message(data)
|
||||
else:
|
||||
logger.warning(f"收到未知类型的消息: {msg_type}")
|
||||
except Exception as e:
|
||||
logger.error(f"处理JSON消息时出错: {e}")
|
||||
|
||||
def _handle_tts_message(self, data):
|
||||
"""处理TTS消息"""
|
||||
state = data.get("state", "")
|
||||
if state == "start":
|
||||
self.schedule(lambda: self._handle_tts_start())
|
||||
elif state == "stop":
|
||||
self.schedule(lambda: self._handle_tts_stop())
|
||||
elif state == "sentence_start":
|
||||
text = data.get("text", "")
|
||||
if text:
|
||||
logger.info(f"<< {text}")
|
||||
self.schedule(lambda: self.set_chat_message("assistant", text))
|
||||
|
||||
# 检查是否包含验证码信息
|
||||
if "请登录到控制面板添加设备,输入验证码" in text:
|
||||
self.schedule(lambda: self._handle_verification_code(text))
|
||||
|
||||
def _handle_tts_start(self):
|
||||
"""处理TTS开始事件"""
|
||||
self.aborted = False
|
||||
|
||||
# 清空可能存在的旧音频数据
|
||||
self.audio_codec.clear_audio_queue()
|
||||
|
||||
if (
|
||||
self.device_state == DeviceState.IDLE
|
||||
or self.device_state == DeviceState.LISTENING
|
||||
):
|
||||
self.set_device_state(DeviceState.SPEAKING)
|
||||
|
||||
def _handle_tts_stop(self):
|
||||
"""处理TTS停止事件"""
|
||||
if self.device_state == DeviceState.SPEAKING:
|
||||
# 给音频播放一个缓冲时间,确保所有音频都播放完毕
|
||||
def delayed_state_change():
|
||||
# 等待音频队列清空
|
||||
self.audio_codec.wait_for_audio_complete()
|
||||
|
||||
# 状态转换
|
||||
if self.keep_listening:
|
||||
asyncio.run_coroutine_threadsafe(
|
||||
self.protocol.send_start_listening(ListeningMode.AUTO_STOP),
|
||||
self.loop,
|
||||
)
|
||||
self.set_device_state(DeviceState.LISTENING)
|
||||
else:
|
||||
self.set_device_state(DeviceState.IDLE)
|
||||
|
||||
# 安排延迟执行
|
||||
threading.Thread(target=delayed_state_change, daemon=True).start()
|
||||
|
||||
def _handle_stt_message(self, data):
|
||||
"""处理STT消息"""
|
||||
text = data.get("text", "")
|
||||
if text:
|
||||
logger.info(f">> {text}")
|
||||
self.schedule(lambda: self.set_chat_message("user", text))
|
||||
|
||||
def _handle_llm_message(self, data):
|
||||
"""处理LLM消息"""
|
||||
emotion = data.get("emotion", "")
|
||||
if emotion:
|
||||
self.schedule(lambda: self.set_emotion(emotion))
|
||||
|
||||
async def _on_audio_channel_opened(self):
|
||||
"""音频通道打开回调"""
|
||||
logger.info("音频通道已打开")
|
||||
self.schedule(lambda: self._start_audio_streams())
|
||||
|
||||
def _start_audio_streams(self):
|
||||
"""启动音频流"""
|
||||
try:
|
||||
# 确保流已关闭后再重新打开
|
||||
if (
|
||||
self.audio_codec.input_stream
|
||||
and self.audio_codec.input_stream.is_active()
|
||||
):
|
||||
self.audio_codec.input_stream.stop_stream()
|
||||
|
||||
# 重新打开流
|
||||
self.audio_codec.input_stream.start_stream()
|
||||
|
||||
if (
|
||||
self.audio_codec.output_stream
|
||||
and self.audio_codec.output_stream.is_active()
|
||||
):
|
||||
self.audio_codec.output_stream.stop_stream()
|
||||
|
||||
# 重新打开流
|
||||
self.audio_codec.output_stream.start_stream()
|
||||
|
||||
# 设置事件触发器
|
||||
threading.Thread(
|
||||
target=self._audio_input_event_trigger, daemon=True
|
||||
).start()
|
||||
threading.Thread(
|
||||
target=self._audio_output_event_trigger, daemon=True
|
||||
).start()
|
||||
|
||||
logger.info("音频流已启动")
|
||||
except Exception as e:
|
||||
logger.error(f"启动音频流失败: {e}")
|
||||
|
||||
def _audio_input_event_trigger(self):
|
||||
"""音频输入事件触发器"""
|
||||
while self.running:
|
||||
try:
|
||||
if (
|
||||
self.audio_codec.input_stream
|
||||
and self.audio_codec.input_stream.is_active()
|
||||
):
|
||||
self.events[EventType.AUDIO_INPUT_READY_EVENT].set()
|
||||
except OSError as e:
|
||||
logger.error(f"音频输入流错误: {e}")
|
||||
# 如果流已关闭,尝试重新打开或者退出循环
|
||||
if "Stream not open" in str(e):
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(f"音频输入事件触发器错误: {e}")
|
||||
|
||||
time.sleep(AudioConfig.FRAME_DURATION / 1000) # 按帧时长触发
|
||||
|
||||
def _audio_output_event_trigger(self):
|
||||
"""音频输出事件触发器"""
|
||||
while (
|
||||
self.running
|
||||
and self.audio_codec.output_stream
|
||||
and self.audio_codec.output_stream.is_active()
|
||||
):
|
||||
# 当队列中有数据时才触发事件
|
||||
if (
|
||||
not self.audio_codec.audio_decode_queue.empty()
|
||||
): # 修改为使用 audio_codec 的队列
|
||||
self.events[EventType.AUDIO_OUTPUT_READY_EVENT].set()
|
||||
time.sleep(0.02) # 稍微延长检查间隔
|
||||
|
||||
async def _on_audio_channel_closed(self):
|
||||
"""音频通道关闭回调"""
|
||||
logger.info("音频通道已关闭")
|
||||
self.set_device_state(DeviceState.IDLE)
|
||||
self.keep_listening = False
|
||||
self.schedule(lambda: self._stop_audio_streams())
|
||||
|
||||
def _stop_audio_streams(self):
|
||||
"""停止音频流"""
|
||||
try:
|
||||
if (
|
||||
self.audio_codec.input_stream
|
||||
and self.audio_codec.input_stream.is_active()
|
||||
):
|
||||
self.audio_codec.input_stream.stop_stream()
|
||||
|
||||
if (
|
||||
self.audio_codec.output_stream
|
||||
and self.audio_codec.output_stream.is_active()
|
||||
):
|
||||
self.audio_codec.output_stream.stop_stream()
|
||||
|
||||
logger.info("音频流已停止")
|
||||
except Exception as e:
|
||||
logger.error(f"停止音频流失败: {e}")
|
||||
|
||||
def set_device_state(self, state):
|
||||
"""设置设备状态"""
|
||||
if self.device_state == state:
|
||||
return
|
||||
|
||||
old_state = self.device_state
|
||||
|
||||
# 如果从 SPEAKING 状态切换出去,确保音频播放完成
|
||||
if old_state == DeviceState.SPEAKING:
|
||||
self.audio_codec.wait_for_audio_complete()
|
||||
|
||||
self.device_state = state
|
||||
logger.info(f"状态变更: {old_state} -> {state}")
|
||||
|
||||
# 根据状态执行相应操作
|
||||
if state == DeviceState.IDLE:
|
||||
self.display.update_status("待命")
|
||||
self.display.update_emotion("😶")
|
||||
# 停止输出流但不关闭它
|
||||
if (
|
||||
self.audio_codec.output_stream
|
||||
and self.audio_codec.output_stream.is_active()
|
||||
):
|
||||
try:
|
||||
self.audio_codec.output_stream.stop_stream()
|
||||
except Exception as e:
|
||||
logger.warning(f"停止输出流时出错: {e}")
|
||||
elif state == DeviceState.CONNECTING:
|
||||
self.display.update_status("连接中...")
|
||||
elif state == DeviceState.LISTENING:
|
||||
self.display.update_status("聆听中...")
|
||||
self.display.update_emotion("🙂")
|
||||
if (
|
||||
self.audio_codec.input_stream
|
||||
and not self.audio_codec.input_stream.is_active()
|
||||
):
|
||||
try:
|
||||
self.audio_codec.input_stream.start_stream()
|
||||
except Exception as e:
|
||||
logger.warning(f"启动输入流时出错: {e}")
|
||||
# 使用 AudioCodec 类中的方法重新初始化
|
||||
self.audio_codec._reinitialize_input_stream()
|
||||
elif state == DeviceState.SPEAKING:
|
||||
self.display.update_status("说话中...")
|
||||
# 确保输出流处于活跃状态
|
||||
if self.audio_codec.output_stream:
|
||||
if not self.audio_codec.output_stream.is_active():
|
||||
try:
|
||||
self.audio_codec.output_stream.start_stream()
|
||||
except Exception as e:
|
||||
logger.warning(f"启动输出流时出错: {e}")
|
||||
# 使用 AudioCodec 类中的方法重新初始化
|
||||
self.audio_codec._reinitialize_output_stream()
|
||||
# 停止输入流
|
||||
if (
|
||||
self.audio_codec.input_stream
|
||||
and self.audio_codec.input_stream.is_active()
|
||||
):
|
||||
try:
|
||||
self.audio_codec.input_stream.stop_stream()
|
||||
except Exception as e:
|
||||
logger.warning(f"停止输入流时出错: {e}")
|
||||
|
||||
# 通知状态变化
|
||||
for callback in self.on_state_changed_callbacks:
|
||||
try:
|
||||
callback(state)
|
||||
except Exception as e:
|
||||
logger.error(f"执行状态变化回调时出错: {e}")
|
||||
|
||||
def _get_status_text(self):
|
||||
"""获取当前状态文本"""
|
||||
states = {
|
||||
DeviceState.IDLE: "待命",
|
||||
DeviceState.CONNECTING: "连接中...",
|
||||
DeviceState.LISTENING: "聆听中...",
|
||||
DeviceState.SPEAKING: "说话中...",
|
||||
}
|
||||
return states.get(self.device_state, "未知")
|
||||
|
||||
def _get_current_text(self):
|
||||
"""获取当前显示文本"""
|
||||
return self.current_text
|
||||
|
||||
def _get_current_emotion(self):
|
||||
"""获取当前表情"""
|
||||
emotions = {
|
||||
"neutral": "😶",
|
||||
"happy": "🙂",
|
||||
"laughing": "😆",
|
||||
"funny": "😂",
|
||||
"sad": "😔",
|
||||
"angry": "😠",
|
||||
"crying": "😭",
|
||||
"loving": "😍",
|
||||
"embarrassed": "😳",
|
||||
"surprised": "😲",
|
||||
"shocked": "😱",
|
||||
"thinking": "🤔",
|
||||
"winking": "😉",
|
||||
"cool": "😎",
|
||||
"relaxed": "😌",
|
||||
"delicious": "🤤",
|
||||
"kissy": "😘",
|
||||
"confident": "😏",
|
||||
"sleepy": "😴",
|
||||
"silly": "😜",
|
||||
"confused": "🙄",
|
||||
}
|
||||
return emotions.get(self.current_emotion, "😶")
|
||||
|
||||
def set_chat_message(self, role, message):
|
||||
"""设置聊天消息"""
|
||||
self.current_text = message
|
||||
# 更新显示
|
||||
if self.display:
|
||||
self.display.update_text(message)
|
||||
|
||||
def set_emotion(self, emotion):
|
||||
"""设置表情"""
|
||||
self.current_emotion = emotion
|
||||
# 更新显示
|
||||
if self.display:
|
||||
self.display.update_emotion(self._get_current_emotion())
|
||||
|
||||
def start_listening(self):
|
||||
"""开始监听"""
|
||||
self.schedule(self._start_listening_impl)
|
||||
|
||||
def _start_listening_impl(self):
|
||||
"""开始监听的实现"""
|
||||
if not self.protocol:
|
||||
logger.error("协议未初始化")
|
||||
return
|
||||
|
||||
self.keep_listening = False
|
||||
|
||||
if self.device_state == DeviceState.IDLE:
|
||||
self.set_device_state(DeviceState.CONNECTING) # 设置设备状态为连接中
|
||||
|
||||
# 尝试打开音频通道
|
||||
if not self.protocol.is_audio_channel_opened():
|
||||
try:
|
||||
# 等待异步操作完成
|
||||
future = asyncio.run_coroutine_threadsafe(
|
||||
self.protocol.open_audio_channel(), self.loop
|
||||
)
|
||||
# 等待操作完成并获取结果
|
||||
success = future.result(timeout=10.0) # 添加超时时间
|
||||
|
||||
if not success:
|
||||
self.alert("错误", "打开音频通道失败") # 弹出错误提示
|
||||
self.set_device_state(DeviceState.IDLE) # 设置设备状态为空闲
|
||||
return
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"打开音频通道时发生错误: {e}")
|
||||
self.alert("错误", f"打开音频通道失败: {str(e)}")
|
||||
self.set_device_state(DeviceState.IDLE)
|
||||
return
|
||||
|
||||
asyncio.run_coroutine_threadsafe(
|
||||
self.protocol.send_start_listening(ListeningMode.MANUAL), self.loop
|
||||
)
|
||||
self.set_device_state(DeviceState.LISTENING) # 设置设备状态为监听中
|
||||
elif self.device_state == DeviceState.SPEAKING:
|
||||
if not self.aborted:
|
||||
self.abort_speaking(AbortReason.WAKE_WORD_DETECTED)
|
||||
|
||||
async def _open_audio_channel_and_start_manual_listening(self):
|
||||
"""打开音频通道并开始手动监听"""
|
||||
if not await self.protocol.open_audio_channel():
|
||||
self.set_device_state(DeviceState.IDLE)
|
||||
self.alert("错误", "打开音频通道失败")
|
||||
return
|
||||
|
||||
await self.protocol.send_start_listening(ListeningMode.MANUAL)
|
||||
self.set_device_state(DeviceState.LISTENING)
|
||||
|
||||
def toggle_chat_state(self):
|
||||
"""切换聊天状态"""
|
||||
self.schedule(self._toggle_chat_state_impl)
|
||||
|
||||
def _toggle_chat_state_impl(self):
|
||||
"""切换聊天状态的具体实现"""
|
||||
# 检查协议是否已初始化
|
||||
if not self.protocol:
|
||||
logger.error("协议未初始化")
|
||||
return
|
||||
|
||||
# 如果设备当前处于空闲状态,尝试连接并开始监听
|
||||
if self.device_state == DeviceState.IDLE:
|
||||
self.set_device_state(DeviceState.CONNECTING) # 设置设备状态为连接中
|
||||
|
||||
# 尝试打开音频通道
|
||||
if not self.protocol.is_audio_channel_opened():
|
||||
try:
|
||||
# 等待异步操作完成
|
||||
future = asyncio.run_coroutine_threadsafe(
|
||||
self.protocol.open_audio_channel(), self.loop
|
||||
)
|
||||
# 等待操作完成并获取结果
|
||||
success = future.result(timeout=10.0) # 添加超时时间
|
||||
|
||||
if not success:
|
||||
self.alert("错误", "打开音频通道失败") # 弹出错误提示
|
||||
self.set_device_state(DeviceState.IDLE) # 设置设备状态为空闲
|
||||
return
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"打开音频通道时发生错误: {e}")
|
||||
self.alert("错误", f"打开音频通道失败: {str(e)}")
|
||||
self.set_device_state(DeviceState.IDLE)
|
||||
return
|
||||
|
||||
self.keep_listening = True # 开始监听
|
||||
# 启动自动停止的监听模式
|
||||
asyncio.run_coroutine_threadsafe(
|
||||
self.protocol.send_start_listening(ListeningMode.AUTO_STOP), self.loop
|
||||
)
|
||||
self.set_device_state(DeviceState.LISTENING) # 设置设备状态为监听中
|
||||
|
||||
# 如果设备正在说话,停止当前说话
|
||||
elif self.device_state == DeviceState.SPEAKING:
|
||||
self.abort_speaking(AbortReason.NONE) # 中止说话
|
||||
|
||||
# 如果设备正在监听,关闭音频通道
|
||||
elif self.device_state == DeviceState.LISTENING:
|
||||
asyncio.run_coroutine_threadsafe(
|
||||
self.protocol.close_audio_channel(), self.loop
|
||||
)
|
||||
|
||||
def stop_listening(self):
|
||||
"""停止监听"""
|
||||
self.schedule(self._stop_listening_impl)
|
||||
|
||||
def _stop_listening_impl(self):
|
||||
"""停止监听的实现"""
|
||||
if self.device_state == DeviceState.LISTENING:
|
||||
asyncio.run_coroutine_threadsafe(
|
||||
self.protocol.send_stop_listening(), self.loop
|
||||
)
|
||||
self.set_device_state(DeviceState.IDLE)
|
||||
|
||||
def abort_speaking(self, reason):
|
||||
"""中止语音输出"""
|
||||
logger.info(f"中止语音输出,原因: {reason}")
|
||||
self.aborted = True
|
||||
asyncio.run_coroutine_threadsafe(
|
||||
self.protocol.send_abort_speaking(reason), self.loop
|
||||
)
|
||||
self.set_device_state(DeviceState.IDLE)
|
||||
|
||||
# 添加此代码:当用户主动打断时自动进入录音模式
|
||||
if reason == AbortReason.WAKE_WORD_DETECTED and self.keep_listening:
|
||||
# 短暂延迟确保abort命令被处理
|
||||
def start_listening_after_abort():
|
||||
time.sleep(0.2) # 短暂延迟
|
||||
self.set_device_state(DeviceState.IDLE)
|
||||
self.schedule(lambda: self.toggle_chat_state())
|
||||
|
||||
threading.Thread(target=start_listening_after_abort, daemon=True).start()
|
||||
|
||||
def alert(self, title, message):
|
||||
"""显示警告信息"""
|
||||
logger.warning(f"警告: {title}, {message}")
|
||||
# 在GUI上显示警告
|
||||
if self.display:
|
||||
self.display.update_text(f"{title}: {message}")
|
||||
|
||||
def on_state_changed(self, callback):
|
||||
"""注册状态变化回调"""
|
||||
self.on_state_changed_callbacks.append(callback)
|
||||
|
||||
def shutdown(self):
|
||||
"""关闭应用程序"""
|
||||
logger.info("正在关闭应用程序...")
|
||||
self.running = False
|
||||
|
||||
# 关闭音频编解码器
|
||||
if self.audio_codec:
|
||||
self.audio_codec.close()
|
||||
|
||||
# 关闭协议
|
||||
if self.protocol:
|
||||
asyncio.run_coroutine_threadsafe(
|
||||
self.protocol.close_audio_channel(), self.loop
|
||||
)
|
||||
|
||||
# 停止事件循环
|
||||
if self.loop and self.loop.is_running():
|
||||
self.loop.call_soon_threadsafe(self.loop.stop)
|
||||
|
||||
# 等待事件循环线程结束
|
||||
if self.loop_thread and self.loop_thread.is_alive():
|
||||
self.loop_thread.join(timeout=1.0)
|
||||
|
||||
logger.info("应用程序已关闭")
|
||||
|
||||
def _handle_verification_code(self, text):
|
||||
"""处理验证码信息"""
|
||||
try:
|
||||
# 提取验证码
|
||||
import re
|
||||
|
||||
verification_code = re.search(r"验证码:(\d+)", text)
|
||||
if verification_code:
|
||||
code = verification_code.group(1)
|
||||
|
||||
# 尝试打开浏览器
|
||||
try:
|
||||
import webbrowser
|
||||
|
||||
if webbrowser.open("https://xiaozhi.me/login"):
|
||||
logger.info("已打开登录页面")
|
||||
else:
|
||||
logger.warning("无法打开浏览器")
|
||||
except Exception as e:
|
||||
logger.warning(f"打开浏览器时出错: {e}")
|
||||
|
||||
# 无论如何都显示验证码
|
||||
self.alert("验证码", f"您的验证码是: {code}")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"处理验证码时出错: {e}")
|
||||
|
||||
def _on_mode_changed(self, auto_mode):
|
||||
"""处理对话模式变更"""
|
||||
# 只有在IDLE状态下才允许切换模式
|
||||
if self.device_state != DeviceState.IDLE:
|
||||
self.alert("提示", "只有在待命状态下才能切换对话模式")
|
||||
return False
|
||||
|
||||
self.keep_listening = auto_mode
|
||||
logger.info(f"对话模式已切换为: {'自动' if auto_mode else '手动'}")
|
||||
return True
|
||||
Reference in New Issue
Block a user