refactor: 重构 Open-XiaoAI Client V2
This commit is contained in:
@@ -0,0 +1,250 @@
|
||||
//! # AudioBus - 音频总线
|
||||
//!
|
||||
//! 核心的音频路由模块,采用发布-订阅模式处理多客户端音频流。
|
||||
//!
|
||||
//! ## 设计理念
|
||||
//!
|
||||
//! ```text
|
||||
//! ┌──────────────────────────────────────────┐
|
||||
//! │ AudioBus │
|
||||
//! │ │
|
||||
//! Client A ──UDP──▶│ ┌─────────┐ ┌──────────────────┐ │
|
||||
//! │ │ Ingress │──────▶ BroadcastChannel │ │──▶ Client B, C...
|
||||
//! Client B ──UDP──▶│ │ Router │ └──────────────────┘ │
|
||||
//! │ └─────────┘ │
|
||||
//! │ │ │
|
||||
//! │ ▼ │
|
||||
//! │ ┌─────────────┐ │
|
||||
//! │ │ Recorder │──▶ WAV File │
|
||||
//! │ └─────────────┘ │
|
||||
//! │ │
|
||||
//! │ ┌─────────────┐ │
|
||||
//! │ │ Playback │◀── WAV/MP3 File │──▶ Client X
|
||||
//! │ └─────────────┘ │
|
||||
//! └──────────────────────────────────────────┘
|
||||
//! ```
|
||||
//!
|
||||
//! ## 核心概念
|
||||
//!
|
||||
//! - **Subscriber**: 订阅者,可以是客户端或录音器
|
||||
//! - **Publisher**: 发布者,可以是客户端麦克风或文件播放器
|
||||
//! - **Channel**: 频道,用于隔离不同的音频流组
|
||||
|
||||
use crate::net::network::AudioSocket;
|
||||
use crate::net::protocol::AudioPacket;
|
||||
use dashmap::DashMap;
|
||||
use std::net::SocketAddr;
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
use tokio::sync::broadcast;
|
||||
|
||||
/// 订阅者 ID
|
||||
pub type SubscriberId = u64;
|
||||
|
||||
/// 音频包及其来源
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct AudioFrame {
|
||||
/// 音频数据包
|
||||
pub packet: AudioPacket,
|
||||
/// 发送者地址(如果来自网络)
|
||||
pub source: Option<SocketAddr>,
|
||||
/// 时间戳(单调递增)
|
||||
pub timestamp: u64,
|
||||
}
|
||||
|
||||
/// 订阅者信息
|
||||
#[derive(Clone)]
|
||||
pub struct Subscriber {
|
||||
pub id: SubscriberId,
|
||||
/// UDP 目标地址
|
||||
pub addr: SocketAddr,
|
||||
/// 是否过滤自己的音频(防止回声)
|
||||
pub filter_self: bool,
|
||||
}
|
||||
|
||||
/// 音频总线 - 负责音频流的路由和分发
|
||||
pub struct AudioBus {
|
||||
/// UDP socket 用于收发音频
|
||||
socket: Arc<AudioSocket>,
|
||||
|
||||
/// 广播通道 - 发布订阅模式的核心
|
||||
/// 所有接收到的音频都会广播到这个 channel
|
||||
broadcast_tx: broadcast::Sender<AudioFrame>,
|
||||
|
||||
/// 订阅者注册表
|
||||
/// key: SocketAddr (UDP 地址)
|
||||
/// value: Subscriber
|
||||
subscribers: DashMap<SocketAddr, Subscriber>,
|
||||
|
||||
/// 地址到订阅者 ID 的反向映射
|
||||
addr_to_id: DashMap<SocketAddr, SubscriberId>,
|
||||
|
||||
/// ID 生成器
|
||||
next_id: AtomicU64,
|
||||
|
||||
/// 全局时间戳
|
||||
timestamp: AtomicU64,
|
||||
}
|
||||
|
||||
impl AudioBus {
|
||||
/// 创建新的音频总线
|
||||
pub async fn new() -> anyhow::Result<Self> {
|
||||
let socket = Arc::new(AudioSocket::bind().await?);
|
||||
// 广播 channel 容量,设置较大以容纳多个订阅者
|
||||
let (broadcast_tx, _) = broadcast::channel(256);
|
||||
|
||||
Ok(Self {
|
||||
socket,
|
||||
broadcast_tx,
|
||||
subscribers: DashMap::new(),
|
||||
addr_to_id: DashMap::new(),
|
||||
next_id: AtomicU64::new(1),
|
||||
timestamp: AtomicU64::new(0),
|
||||
})
|
||||
}
|
||||
|
||||
/// 获取 UDP 端口
|
||||
pub fn port(&self) -> u16 {
|
||||
self.socket.port()
|
||||
}
|
||||
|
||||
/// 获取 socket 引用(用于外部发送)
|
||||
pub fn socket(&self) -> Arc<AudioSocket> {
|
||||
self.socket.clone()
|
||||
}
|
||||
|
||||
/// 注册一个订阅者
|
||||
pub fn register(&self, addr: SocketAddr, filter_self: bool) -> SubscriberId {
|
||||
let id = self.next_id.fetch_add(1, Ordering::SeqCst);
|
||||
let subscriber = Subscriber {
|
||||
id,
|
||||
addr,
|
||||
filter_self,
|
||||
};
|
||||
self.subscribers.insert(addr, subscriber);
|
||||
self.addr_to_id.insert(addr, id);
|
||||
println!("[AudioBus] Subscriber {id} registered at {addr}");
|
||||
id
|
||||
}
|
||||
|
||||
/// 注销订阅者
|
||||
pub fn unregister(&self, addr: &SocketAddr) {
|
||||
if let Some((_, sub)) = self.subscribers.remove(addr) {
|
||||
self.addr_to_id.remove(addr);
|
||||
println!("[AudioBus] Subscriber {} unregistered", sub.id);
|
||||
}
|
||||
}
|
||||
|
||||
/// 订阅广播频道,返回一个 Receiver
|
||||
pub fn subscribe(&self) -> broadcast::Receiver<AudioFrame> {
|
||||
self.broadcast_tx.subscribe()
|
||||
}
|
||||
|
||||
/// 发布音频帧到总线
|
||||
pub fn publish(&self, packet: AudioPacket, source: Option<SocketAddr>) {
|
||||
let ts = self.timestamp.fetch_add(1, Ordering::SeqCst);
|
||||
let frame = AudioFrame {
|
||||
packet,
|
||||
source,
|
||||
timestamp: ts,
|
||||
};
|
||||
// 忽略没有订阅者的情况
|
||||
let _ = self.broadcast_tx.send(frame);
|
||||
}
|
||||
|
||||
/// 广播音频帧到所有订阅者(除了发送者自己)
|
||||
pub async fn broadcast(&self, frame: &AudioFrame) {
|
||||
for entry in self.subscribers.iter() {
|
||||
let sub = entry.value();
|
||||
|
||||
// 如果启用了自过滤,跳过发送者自己
|
||||
if sub.filter_self {
|
||||
if let Some(src) = &frame.source {
|
||||
if src == &sub.addr {
|
||||
continue;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 发送音频包
|
||||
if let Err(e) = self.socket.send(&frame.packet, sub.addr).await {
|
||||
eprintln!("[AudioBus] Failed to send to {}: {}", sub.addr, e);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 发送音频包到指定地址
|
||||
pub async fn send_to(&self, packet: &AudioPacket, addr: SocketAddr) -> anyhow::Result<()> {
|
||||
self.socket.send(packet, addr).await
|
||||
}
|
||||
|
||||
/// 启动 UDP 接收循环
|
||||
/// 这是一个独立的任务,负责:
|
||||
/// 1. 接收 UDP 音频包
|
||||
/// 2. 发布到广播频道
|
||||
pub async fn run_receiver(&self) {
|
||||
let mut buf = vec![0u8; 4096];
|
||||
loop {
|
||||
match self.socket.recv(&mut buf).await {
|
||||
Ok((packet, src_addr)) => {
|
||||
// 只处理已注册的发送者
|
||||
if self.subscribers.contains_key(&src_addr) {
|
||||
self.publish(packet, Some(src_addr));
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
eprintln!("[AudioBus] Recv error: {}", e);
|
||||
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 启动广播分发循环
|
||||
/// 订阅广播频道并将音频转发给所有订阅者
|
||||
pub async fn run_broadcaster(self: Arc<Self>) {
|
||||
let mut rx = self.subscribe();
|
||||
loop {
|
||||
match rx.recv().await {
|
||||
Ok(frame) => {
|
||||
self.broadcast(&frame).await;
|
||||
}
|
||||
Err(broadcast::error::RecvError::Lagged(n)) => {
|
||||
eprintln!("[AudioBus] Broadcaster lagged {} frames", n);
|
||||
}
|
||||
Err(broadcast::error::RecvError::Closed) => {
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取当前订阅者数量
|
||||
pub fn subscriber_count(&self) -> usize {
|
||||
self.subscribers.len()
|
||||
}
|
||||
|
||||
/// 检查地址是否已注册
|
||||
pub fn is_registered(&self, addr: &SocketAddr) -> bool {
|
||||
self.subscribers.contains_key(addr)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_audio_bus_basic() {
|
||||
let bus = AudioBus::new().await.unwrap();
|
||||
let addr: SocketAddr = "127.0.0.1:12345".parse().unwrap();
|
||||
|
||||
let id = bus.register(addr, true);
|
||||
assert!(bus.is_registered(&addr));
|
||||
assert_eq!(bus.subscriber_count(), 1);
|
||||
|
||||
bus.unregister(&addr);
|
||||
assert!(!bus.is_registered(&addr));
|
||||
assert_eq!(bus.subscriber_count(), 0);
|
||||
}
|
||||
}
|
||||
@@ -1,208 +0,0 @@
|
||||
use crate::audio::codec::OpusCodec;
|
||||
use crate::audio::config::AudioConfig;
|
||||
use crate::audio::wav::{WavReader, WavWriter};
|
||||
use crate::net::network::AudioSocket;
|
||||
use crate::net::protocol::AudioPacket;
|
||||
use anyhow::Result;
|
||||
use parking_lot::Mutex;
|
||||
use std::net::SocketAddr;
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::mpsc;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
use tokio_util::task::TaskTracker;
|
||||
|
||||
pub enum RecorderCommand {
|
||||
Start {
|
||||
config: AudioConfig,
|
||||
filename: String,
|
||||
},
|
||||
Stop,
|
||||
}
|
||||
|
||||
pub struct ServerAudioManager {
|
||||
session_cancel: CancellationToken,
|
||||
record_cancel: Mutex<Option<CancellationToken>>,
|
||||
play_cancel: Mutex<Option<CancellationToken>>,
|
||||
audio_tx: mpsc::Sender<AudioPacket>,
|
||||
recorder_tx: mpsc::Sender<RecorderCommand>,
|
||||
tracker: TaskTracker,
|
||||
}
|
||||
|
||||
impl ServerAudioManager {
|
||||
pub fn new(
|
||||
session_cancel: CancellationToken,
|
||||
tracker: TaskTracker,
|
||||
) -> (
|
||||
Self,
|
||||
mpsc::Receiver<AudioPacket>,
|
||||
mpsc::Receiver<RecorderCommand>,
|
||||
) {
|
||||
let (audio_tx, audio_rx) = mpsc::channel(1024);
|
||||
let (recorder_tx, recorder_rx) = mpsc::channel(64);
|
||||
|
||||
let manager = Self {
|
||||
session_cancel,
|
||||
record_cancel: Mutex::new(None),
|
||||
play_cancel: Mutex::new(None),
|
||||
audio_tx,
|
||||
recorder_tx,
|
||||
tracker,
|
||||
};
|
||||
|
||||
(manager, audio_rx, recorder_rx)
|
||||
}
|
||||
|
||||
pub fn audio_tx(&self) -> mpsc::Sender<AudioPacket> {
|
||||
self.audio_tx.clone()
|
||||
}
|
||||
|
||||
pub async fn start_recording(&self, config: AudioConfig, filename: String) -> Result<()> {
|
||||
// 仅停止之前的录音任务
|
||||
{
|
||||
let mut guard = self.record_cancel.lock();
|
||||
if let Some(token) = guard.take() {
|
||||
token.cancel();
|
||||
}
|
||||
*guard = Some(self.session_cancel.child_token());
|
||||
}
|
||||
|
||||
self.recorder_tx
|
||||
.send(RecorderCommand::Start { config, filename })
|
||||
.await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn stop_recording(&self) -> Result<()> {
|
||||
if let Some(token) = self.record_cancel.lock().take() {
|
||||
token.cancel();
|
||||
}
|
||||
let _ = self.recorder_tx.send(RecorderCommand::Stop).await;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn start_playback(
|
||||
&self,
|
||||
config: AudioConfig,
|
||||
reader: WavReader,
|
||||
audio_socket: Arc<AudioSocket>,
|
||||
target_addr: SocketAddr,
|
||||
) -> Result<()> {
|
||||
let token = {
|
||||
let mut guard = self.play_cancel.lock();
|
||||
if let Some(token) = guard.take() {
|
||||
token.cancel();
|
||||
}
|
||||
let token = self.session_cancel.child_token();
|
||||
*guard = Some(token.clone());
|
||||
token
|
||||
};
|
||||
|
||||
let session_cancel = self.session_cancel.clone();
|
||||
|
||||
self.tracker.spawn(async move {
|
||||
let mut reader = reader;
|
||||
let mut codec = match OpusCodec::new(&config) {
|
||||
Ok(c) => c,
|
||||
Err(e) => {
|
||||
eprintln!("Failed to create opus codec: {}", e);
|
||||
return;
|
||||
}
|
||||
};
|
||||
let mut pcm = vec![0i16; config.frame_size];
|
||||
let mut opus = vec![0u8; 4096];
|
||||
let mut interval = tokio::time::interval(std::time::Duration::from_millis(20));
|
||||
|
||||
loop {
|
||||
tokio::select! {
|
||||
_ = token.cancelled() => break,
|
||||
_ = session_cancel.cancelled() => break,
|
||||
_ = interval.tick() => {
|
||||
match reader.read_samples(&mut pcm) {
|
||||
Ok(0) => break,
|
||||
Ok(n) => {
|
||||
if let Ok(len) = codec.encode(&pcm[..n], &mut opus) {
|
||||
let _ = audio_socket
|
||||
.send(
|
||||
&AudioPacket {
|
||||
data: opus[..len].to_vec(),
|
||||
},
|
||||
target_addr,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
eprintln!("Failed to read samples: {}", e);
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn stop_playback(&self) {
|
||||
if let Some(token) = self.play_cancel.lock().take() {
|
||||
token.cancel();
|
||||
}
|
||||
}
|
||||
|
||||
pub fn spawn_audio_processor(
|
||||
&self,
|
||||
mut audio_rx: mpsc::Receiver<AudioPacket>,
|
||||
mut recorder_rx: mpsc::Receiver<RecorderCommand>,
|
||||
) {
|
||||
let session_cancel = self.session_cancel.clone();
|
||||
self.tracker.spawn(async move {
|
||||
let mut active_recorder: Option<(WavWriter, OpusCodec, usize)> = None;
|
||||
loop {
|
||||
tokio::select! {
|
||||
_ = session_cancel.cancelled() => break,
|
||||
cmd = recorder_rx.recv() => {
|
||||
match cmd {
|
||||
Some(RecorderCommand::Start { config, filename }) => {
|
||||
if let Some((writer, _, _)) = active_recorder.take() {
|
||||
let _ = writer.finalize();
|
||||
}
|
||||
match WavWriter::create(&filename, config.sample_rate, config.channels) {
|
||||
Ok(writer) => {
|
||||
match OpusCodec::new(&config) {
|
||||
Ok(codec) => active_recorder = Some((writer, codec, config.frame_size)),
|
||||
Err(e) => eprintln!("Failed to create opus codec: {}", e),
|
||||
}
|
||||
}
|
||||
Err(e) => eprintln!("Failed to create wav writer: {}", e),
|
||||
}
|
||||
}
|
||||
Some(RecorderCommand::Stop) => {
|
||||
if let Some((writer, _, _)) = active_recorder.take() {
|
||||
let _ = writer.finalize();
|
||||
}
|
||||
}
|
||||
None => break,
|
||||
}
|
||||
}
|
||||
packet = audio_rx.recv() => {
|
||||
match packet {
|
||||
Some(packet) => {
|
||||
if let Some((writer, codec, frame_size)) = &mut active_recorder {
|
||||
let mut pcm = vec![0i16; *frame_size];
|
||||
if let Ok(n) = codec.decode(&packet.data, &mut pcm) {
|
||||
let _ = writer.write_samples(&pcm[..n]);
|
||||
}
|
||||
}
|
||||
}
|
||||
None => break,
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
if let Some((writer, _, _)) = active_recorder.take() {
|
||||
let _ = writer.finalize();
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
@@ -1,109 +1,228 @@
|
||||
mod audio_manager;
|
||||
//! # Server 模块
|
||||
//!
|
||||
//! 实时音频流服务器,支持:
|
||||
//! - 多客户端连接管理
|
||||
//! - 音频流发布-订阅
|
||||
//! - 录音和播放
|
||||
//! - RPC 远程调用
|
||||
//! - 实时事件推送
|
||||
//!
|
||||
//! ## 架构
|
||||
//!
|
||||
//! ```text
|
||||
//! ┌─────────────────────────────────────────────┐
|
||||
//! │ Server │
|
||||
//! │ │
|
||||
//! Client 1 ──TCP────────┼──▶ SessionManager │
|
||||
//! Client 2 ──TCP────────┼──▶ ├─ Session 1 │
|
||||
//! Client N ──TCP────────┼──▶ ├─ Session 2 │
|
||||
//! │ └─ Session N │
|
||||
//! │ │ │
|
||||
//! ┌────────────────┼────────────┴────────────────────────┐ │
|
||||
//! │ │ │ │
|
||||
//! ▼ │ ▼ │
|
||||
//! ┌─────────┐ │ ┌─────────────────────────────┐ │
|
||||
//! │Command │ │ │ AudioBus │ │
|
||||
//! │Handler │ │ │ │ │
|
||||
//! └─────────┘ │ │ ┌─────────┐ ┌──────────┐ │ │
|
||||
//! │ │ │ │Receiver │ │Broadcaster│ │ │
|
||||
//! ▼ │ │ │ Loop │─▶│ Loop │──┼───┼──▶ All Clients
|
||||
//! ┌─────────┐ │ │ └─────────┘ └──────────┘ │ │
|
||||
//! │ Event │ │ └─────────────────────────────┘ │
|
||||
//! │ Bus │───────────┼─────────────────────────────────────────────┼──▶ Events
|
||||
//! └─────────┘ │ │
|
||||
//! └─────────────────────────────────────────────┘
|
||||
//! ```
|
||||
|
||||
mod audio_bus;
|
||||
mod session;
|
||||
mod stream;
|
||||
|
||||
pub use audio_bus::{AudioBus, AudioFrame};
|
||||
pub use session::{Session, SessionManager};
|
||||
pub use stream::{FilePlaybackStream, RecorderStream, StreamHandle};
|
||||
|
||||
use crate::audio::config::AudioConfig;
|
||||
use crate::audio::wav::WavReader;
|
||||
use crate::net::command::{
|
||||
AudioState, Command, CommandError, CommandResult, DeviceInfo, ShellResponse,
|
||||
};
|
||||
use crate::net::discovery::Discovery;
|
||||
use crate::net::network::{AudioSocket, Connection};
|
||||
use crate::net::protocol::{ClientInfo, ControlPacket, RpcResult};
|
||||
use crate::net::rpc::RpcManager;
|
||||
use crate::net::event::{ClientEvent, ServerEvent, ServerEventBus};
|
||||
use crate::net::network::Connection;
|
||||
use crate::net::protocol::ControlPacket;
|
||||
use anyhow::{Context, Result, anyhow};
|
||||
use audio_manager::ServerAudioManager;
|
||||
use dashmap::DashMap;
|
||||
use std::net::SocketAddr;
|
||||
use std::sync::Arc;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
use tokio_util::task::TaskTracker;
|
||||
|
||||
pub struct Session {
|
||||
pub info: ClientInfo,
|
||||
pub conn: Arc<Connection>,
|
||||
pub rpc: Arc<RpcManager>,
|
||||
pub tcp_addr: SocketAddr,
|
||||
pub audio_addr: SocketAddr,
|
||||
pub session_cancel: CancellationToken,
|
||||
pub audio_manager: ServerAudioManager,
|
||||
pub tracker: TaskTracker,
|
||||
}
|
||||
|
||||
/// 实时音频服务器
|
||||
pub struct Server {
|
||||
sessions: DashMap<SocketAddr, Arc<Session>>,
|
||||
udp_to_tcp: DashMap<SocketAddr, SocketAddr>,
|
||||
audio: Arc<AudioSocket>,
|
||||
/// 会话管理器
|
||||
sessions: Arc<SessionManager>,
|
||||
/// 音频总线
|
||||
audio_bus: Arc<AudioBus>,
|
||||
/// 服务端事件总线
|
||||
event_bus: Arc<ServerEventBus>,
|
||||
/// 服务器取消令牌
|
||||
cancel: CancellationToken,
|
||||
/// 服务器启动时间
|
||||
started_at: std::time::Instant,
|
||||
}
|
||||
|
||||
impl Server {
|
||||
/// 创建新服务器
|
||||
pub async fn new() -> Result<Self> {
|
||||
let audio_bus = Arc::new(AudioBus::new().await?);
|
||||
|
||||
Ok(Self {
|
||||
sessions: DashMap::new(),
|
||||
udp_to_tcp: DashMap::new(),
|
||||
audio: Arc::new(AudioSocket::bind().await?),
|
||||
sessions: Arc::new(SessionManager::new()),
|
||||
audio_bus,
|
||||
event_bus: Arc::new(ServerEventBus::default()),
|
||||
cancel: CancellationToken::new(),
|
||||
started_at: std::time::Instant::now(),
|
||||
})
|
||||
}
|
||||
|
||||
/// 获取事件总线(用于外部订阅)
|
||||
pub fn event_bus(&self) -> Arc<ServerEventBus> {
|
||||
self.event_bus.clone()
|
||||
}
|
||||
|
||||
/// 启动服务器
|
||||
pub async fn run(self: Arc<Self>, port: u16) -> Result<()> {
|
||||
let listener = tokio::net::TcpListener::bind(format!("0.0.0.0:{}", port)).await?;
|
||||
let addr = listener.local_addr()?;
|
||||
println!("Server listening on TCP: {}", addr);
|
||||
println!("[Server] Listening on TCP: {}", addr);
|
||||
println!("[Server] Audio UDP port: {}", self.audio_bus.port());
|
||||
|
||||
// 广播服务发现
|
||||
Discovery::broadcast(port).await?;
|
||||
|
||||
// Audio Dispatcher
|
||||
let this = self.clone();
|
||||
// 启动音频总线接收循环
|
||||
let bus = self.audio_bus.clone();
|
||||
tokio::spawn(async move {
|
||||
let mut buf = vec![0u8; 4096];
|
||||
loop {
|
||||
match this.audio.recv(&mut buf).await {
|
||||
Ok((packet, src_addr)) => {
|
||||
if let Some(tcp_addr) = this.udp_to_tcp.get(&src_addr) {
|
||||
if let Some(session) = this.sessions.get(tcp_addr.value()) {
|
||||
// 使用 try_send 避免某一个客户端阻塞导致全局音频延迟
|
||||
let _ = session.audio_manager.audio_tx().try_send(packet);
|
||||
}
|
||||
bus.run_receiver().await;
|
||||
});
|
||||
|
||||
// 启动音频总线广播循环
|
||||
let bus = self.audio_bus.clone();
|
||||
tokio::spawn(async move {
|
||||
bus.run_broadcaster().await;
|
||||
});
|
||||
|
||||
// TCP 连接接受循环
|
||||
loop {
|
||||
tokio::select! {
|
||||
_ = self.cancel.cancelled() => {
|
||||
println!("[Server] Shutting down...");
|
||||
break;
|
||||
}
|
||||
result = listener.accept() => {
|
||||
match result {
|
||||
Ok((stream, addr)) => {
|
||||
let server = self.clone();
|
||||
tokio::spawn(async move {
|
||||
if let Err(e) = server.handle_connection(stream, addr).await {
|
||||
eprintln!("[Server] Connection {} error: {}", addr, e);
|
||||
}
|
||||
});
|
||||
}
|
||||
Err(e) => {
|
||||
eprintln!("[Server] Accept error: {}", e);
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
eprintln!("Audio socket recv error: {}", e);
|
||||
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
loop {
|
||||
let (stream, addr) = listener.accept().await?;
|
||||
let server = self.clone();
|
||||
tokio::spawn(async move {
|
||||
if let Err(e) = server.clone().handle_session(stream, addr).await {
|
||||
eprintln!("Session {} error: {}", addr, e);
|
||||
}
|
||||
server.remove_session(&addr).await;
|
||||
});
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn remove_session(&self, tcp_addr: &SocketAddr) {
|
||||
if let Some((_, session)) = self.sessions.remove(tcp_addr) {
|
||||
session.session_cancel.cancel();
|
||||
session.tracker.close();
|
||||
self.udp_to_tcp.remove(&session.audio_addr);
|
||||
println!("Session {} ({}) closed", tcp_addr, session.info.model);
|
||||
}
|
||||
}
|
||||
|
||||
async fn handle_session(
|
||||
/// 处理新连接
|
||||
async fn handle_connection(
|
||||
self: Arc<Self>,
|
||||
stream: tokio::net::TcpStream,
|
||||
addr: SocketAddr,
|
||||
) -> Result<()> {
|
||||
println!("New TCP connection from {}", addr);
|
||||
println!("[Server] New connection from {}", addr);
|
||||
let conn = Arc::new(Connection::new(stream)?);
|
||||
|
||||
// --- 握手 (Handshake) ---
|
||||
// --- 握手 ---
|
||||
let (info, audio_addr) = self.handshake(&conn, addr).await?;
|
||||
println!(
|
||||
"[Server] Client identified: {} ({}) audio: {}",
|
||||
info.model, info.serial_number, audio_addr
|
||||
);
|
||||
|
||||
// --- 创建 Session ---
|
||||
let session_cancel = self.cancel.child_token();
|
||||
let session = Arc::new(Session::new(
|
||||
info.clone(),
|
||||
conn.clone(),
|
||||
addr,
|
||||
audio_addr,
|
||||
session_cancel.clone(),
|
||||
));
|
||||
|
||||
// 注册到 SessionManager 和 AudioBus
|
||||
self.sessions.register(session.clone());
|
||||
self.audio_bus.register(audio_addr, true);
|
||||
|
||||
// 发布客户端加入事件
|
||||
self.event_bus.publish(ServerEvent::ClientJoined {
|
||||
addr: addr.to_string(),
|
||||
model: info.model.clone(),
|
||||
});
|
||||
|
||||
// 广播给其他客户端
|
||||
self.sessions
|
||||
.broadcast_except(
|
||||
&ControlPacket::ServerEvent(ServerEvent::ClientJoined {
|
||||
addr: addr.to_string(),
|
||||
model: info.model.clone(),
|
||||
}),
|
||||
&addr,
|
||||
)
|
||||
.await;
|
||||
|
||||
// --- 主循环 ---
|
||||
let result = self.session_loop(session.clone()).await;
|
||||
|
||||
// --- 清理 ---
|
||||
self.audio_bus.unregister(&audio_addr);
|
||||
self.sessions.unregister(&addr);
|
||||
|
||||
// 发布客户端离开事件
|
||||
self.event_bus.publish(ServerEvent::ClientLeft {
|
||||
addr: addr.to_string(),
|
||||
model: info.model.clone(),
|
||||
});
|
||||
|
||||
// 广播给其他客户端
|
||||
self.sessions
|
||||
.broadcast(&ControlPacket::ServerEvent(ServerEvent::ClientLeft {
|
||||
addr: addr.to_string(),
|
||||
model: info.model,
|
||||
}))
|
||||
.await;
|
||||
|
||||
result
|
||||
}
|
||||
|
||||
/// 握手流程
|
||||
async fn handshake(
|
||||
&self,
|
||||
conn: &Arc<Connection>,
|
||||
addr: SocketAddr,
|
||||
) -> Result<(crate::net::protocol::ClientInfo, SocketAddr)> {
|
||||
let version = env!("CARGO_PKG_VERSION").to_string();
|
||||
let server_auth =
|
||||
std::env::var("XIAO_SERVER_AUTH").unwrap_or_else(|_| "xiao-server".to_string());
|
||||
let client_auth =
|
||||
std::env::var("XIAO_CLIENT_AUTH").unwrap_or_else(|_| "xiao-client".to_string());
|
||||
|
||||
// 等待客户端 Hello
|
||||
let (info, client_audio_port) = match conn.recv().await? {
|
||||
ControlPacket::ClientHello {
|
||||
auth,
|
||||
@@ -112,71 +231,45 @@ impl Server {
|
||||
info,
|
||||
} => {
|
||||
if v != version {
|
||||
return Err(anyhow!("Client version mismatch: {} != {}", v, version));
|
||||
return Err(anyhow!("Version mismatch: {} != {}", v, version));
|
||||
}
|
||||
if auth != server_auth {
|
||||
return Err(anyhow!("Invalid client auth"));
|
||||
}
|
||||
(info, udp_port)
|
||||
}
|
||||
_ => return Err(anyhow!("Handshake failed")),
|
||||
_ => return Err(anyhow!("Expected ClientHello")),
|
||||
};
|
||||
|
||||
// 发送服务器 Hello
|
||||
conn.send(&ControlPacket::ServerHello {
|
||||
auth: client_auth,
|
||||
version: env!("CARGO_PKG_VERSION").to_string(),
|
||||
udp_port: self.audio.port(),
|
||||
version,
|
||||
udp_port: self.audio_bus.port(),
|
||||
})
|
||||
.await?;
|
||||
|
||||
let audio_addr = SocketAddr::new(addr.ip(), client_audio_port);
|
||||
println!(
|
||||
"Client identified: {} ({}), audio at {}",
|
||||
info.model, info.serial_number, audio_addr
|
||||
);
|
||||
Ok((info, audio_addr))
|
||||
}
|
||||
|
||||
// --- 初始化 Session ---
|
||||
let tracker = TaskTracker::new();
|
||||
let session_cancel = CancellationToken::new();
|
||||
let (audio_manager, audio_rx, recorder_rx) =
|
||||
ServerAudioManager::new(session_cancel.clone(), tracker.clone());
|
||||
/// Session 消息循环
|
||||
async fn session_loop(&self, session: Arc<Session>) -> Result<()> {
|
||||
let timeout = std::time::Duration::from_secs(60);
|
||||
|
||||
let session = Arc::new(Session {
|
||||
info,
|
||||
conn: conn.clone(),
|
||||
rpc: Arc::new(RpcManager::new()),
|
||||
tcp_addr: addr,
|
||||
audio_addr,
|
||||
session_cancel,
|
||||
audio_manager,
|
||||
tracker: tracker.clone(),
|
||||
});
|
||||
|
||||
session
|
||||
.audio_manager
|
||||
.spawn_audio_processor(audio_rx, recorder_rx);
|
||||
|
||||
self.sessions.insert(addr, session.clone());
|
||||
self.udp_to_tcp.insert(audio_addr, addr);
|
||||
|
||||
// 消息主循环
|
||||
loop {
|
||||
tokio::select! {
|
||||
_ = session.session_cancel.cancelled() => break,
|
||||
res = tokio::time::timeout(std::time::Duration::from_secs(60), conn.recv()) => {
|
||||
match res {
|
||||
_ = session.cancel.cancelled() => break,
|
||||
result = tokio::time::timeout(timeout, session.recv()) => {
|
||||
match result {
|
||||
Ok(Ok(packet)) => {
|
||||
if let Err(e) = self.process_packet(session.clone(), packet).await {
|
||||
eprintln!("Process packet error: {}", e);
|
||||
}
|
||||
self.handle_packet(&session, packet).await?;
|
||||
}
|
||||
Ok(Err(e)) => {
|
||||
eprintln!("Session {} connection error: {}", addr, e);
|
||||
break;
|
||||
return Err(anyhow!("Connection error: {}", e));
|
||||
}
|
||||
Err(_) => {
|
||||
eprintln!("Session {} timeout", addr);
|
||||
break;
|
||||
return Err(anyhow!("Connection timeout"));
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -186,119 +279,209 @@ impl Server {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn get_clients(&self) -> Vec<SocketAddr> {
|
||||
self.sessions.iter().map(|r| *r.key()).collect()
|
||||
}
|
||||
|
||||
// todo 考虑有些操作比较耗时,需要非阻塞处理
|
||||
async fn process_packet(&self, session: Arc<Session>, packet: ControlPacket) -> Result<()> {
|
||||
/// 处理控制包
|
||||
async fn handle_packet(&self, session: &Arc<Session>, packet: ControlPacket) -> Result<()> {
|
||||
match packet {
|
||||
ControlPacket::Ping => {
|
||||
session.conn.send(&ControlPacket::Pong).await?;
|
||||
session.send(&ControlPacket::Pong).await?;
|
||||
}
|
||||
ControlPacket::Pong => {}
|
||||
ControlPacket::RpcResponse { id, result } => {
|
||||
session.rpc.resolve(id, result);
|
||||
session.resolve_rpc(id, result);
|
||||
}
|
||||
ControlPacket::RpcRequest { id, method, args } => {
|
||||
let result = self.handle_rpc(&session, &method, args).await;
|
||||
ControlPacket::RpcRequest { id, command } => {
|
||||
let result = self.handle_command(session, command).await;
|
||||
session
|
||||
.conn
|
||||
.send(&ControlPacket::RpcResponse { id, result })
|
||||
.await?;
|
||||
}
|
||||
ControlPacket::ClientEvent(event) => {
|
||||
self.handle_client_event(session, event).await;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn handle_rpc(
|
||||
&self,
|
||||
_session: &Arc<Session>,
|
||||
method: &str,
|
||||
args: Vec<String>,
|
||||
) -> RpcResult {
|
||||
match method {
|
||||
"hello" => RpcResult {
|
||||
stdout: format!("Hello from server! Args: {:?}", args),
|
||||
..Default::default()
|
||||
},
|
||||
_ => RpcResult {
|
||||
stderr: format!("Unknown method: {}", method),
|
||||
code: -1,
|
||||
..Default::default()
|
||||
},
|
||||
/// 处理 RPC 命令
|
||||
async fn handle_command(&self, session: &Arc<Session>, command: Command) -> CommandResult {
|
||||
match command {
|
||||
Command::Shell(req) => {
|
||||
// 转发给客户端执行
|
||||
match session.execute(Command::Shell(req)).await {
|
||||
Ok(result) => result,
|
||||
Err(e) => CommandResult::Error(CommandError::internal(e.to_string())),
|
||||
}
|
||||
}
|
||||
Command::GetInfo => {
|
||||
// 返回服务器信息
|
||||
CommandResult::Info(DeviceInfo {
|
||||
model: "XiaoAi-Server".to_string(),
|
||||
serial_number: "SERVER-001".to_string(),
|
||||
version: env!("CARGO_PKG_VERSION").to_string(),
|
||||
uptime_secs: self.started_at.elapsed().as_secs(),
|
||||
audio_state: AudioState {
|
||||
is_recording: session.is_recording(),
|
||||
is_playing: session.is_playing(),
|
||||
volume: 100,
|
||||
},
|
||||
})
|
||||
}
|
||||
Command::Ping { timestamp } => {
|
||||
let now = std::time::SystemTime::now()
|
||||
.duration_since(std::time::UNIX_EPOCH)
|
||||
.unwrap()
|
||||
.as_millis() as u64;
|
||||
CommandResult::Pong {
|
||||
timestamp,
|
||||
server_time: now,
|
||||
}
|
||||
}
|
||||
Command::SetVolume(_) => {
|
||||
// 转发给客户端
|
||||
match session.execute(command).await {
|
||||
Ok(result) => result,
|
||||
Err(e) => CommandResult::Error(CommandError::internal(e.to_string())),
|
||||
}
|
||||
}
|
||||
_ => CommandResult::Error(CommandError::not_implemented()),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn call(
|
||||
&self,
|
||||
addr: SocketAddr,
|
||||
method: &str,
|
||||
args: Vec<String>,
|
||||
) -> Result<RpcResult> {
|
||||
let session = self
|
||||
.sessions
|
||||
.get(&addr)
|
||||
.map(|r| r.value().clone())
|
||||
.context("Session not found")?;
|
||||
let (id, rx) = session.rpc.register();
|
||||
session
|
||||
.conn
|
||||
.send(&ControlPacket::RpcRequest {
|
||||
id,
|
||||
method: method.to_string(),
|
||||
args,
|
||||
})
|
||||
.await?;
|
||||
Ok(rx.await?)
|
||||
/// 处理客户端事件
|
||||
async fn handle_client_event(&self, session: &Arc<Session>, event: ClientEvent) {
|
||||
match &event {
|
||||
ClientEvent::Alert { level, message } => {
|
||||
println!(
|
||||
"[Event] Alert from {}: [{:?}] {}",
|
||||
session.tcp_addr, level, message
|
||||
);
|
||||
}
|
||||
ClientEvent::AudioLevel {
|
||||
level_db,
|
||||
is_silent,
|
||||
} => {
|
||||
println!(
|
||||
"[Event] Audio level from {}: {:.1}dB (silent: {})",
|
||||
session.tcp_addr, level_db, is_silent
|
||||
);
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
|
||||
// 可以在这里将客户端事件转发给其他订阅者
|
||||
}
|
||||
|
||||
// ==================== 公开 API ====================
|
||||
|
||||
/// 获取所有已连接的客户端地址
|
||||
pub async fn get_clients(&self) -> Vec<SocketAddr> {
|
||||
self.sessions.all_addrs()
|
||||
}
|
||||
|
||||
/// 获取客户端数量
|
||||
pub fn client_count(&self) -> usize {
|
||||
self.sessions.count()
|
||||
}
|
||||
|
||||
/// 获取服务器运行时间
|
||||
pub fn uptime_secs(&self) -> u64 {
|
||||
self.started_at.elapsed().as_secs()
|
||||
}
|
||||
|
||||
/// 向客户端发起 RPC 调用(新版)
|
||||
pub async fn execute(&self, addr: SocketAddr, command: Command) -> Result<CommandResult> {
|
||||
let session = self.sessions.get(&addr).context("Session not found")?;
|
||||
session.execute(command).await
|
||||
}
|
||||
|
||||
/// 执行 Shell 命令
|
||||
pub async fn shell(&self, addr: SocketAddr, cmd: &str) -> Result<ShellResponse> {
|
||||
let result = self.execute(addr, Command::shell(cmd)).await?;
|
||||
match result {
|
||||
CommandResult::Shell(resp) => Ok(resp),
|
||||
CommandResult::Error(e) => Err(anyhow!("{}", e)),
|
||||
_ => Err(anyhow!("Unexpected response type")),
|
||||
}
|
||||
}
|
||||
|
||||
/// 推送事件给所有客户端
|
||||
pub async fn broadcast_event(&self, event: ServerEvent) {
|
||||
self.event_bus.publish(event.clone());
|
||||
self.sessions
|
||||
.broadcast(&ControlPacket::ServerEvent(event))
|
||||
.await;
|
||||
}
|
||||
|
||||
/// 推送事件给指定客户端
|
||||
pub async fn send_event(&self, addr: SocketAddr, event: ServerEvent) -> Result<()> {
|
||||
let session = self.sessions.get(&addr).context("Session not found")?;
|
||||
self.event_bus.publish_to(addr, event.clone());
|
||||
session.send(&ControlPacket::ServerEvent(event)).await
|
||||
}
|
||||
|
||||
/// 开始录音
|
||||
pub async fn start_record(&self, addr: SocketAddr, config: AudioConfig) -> Result<()> {
|
||||
let session = self
|
||||
.sessions
|
||||
.get(&addr)
|
||||
.map(|r| r.value().clone())
|
||||
.context("Session not found")?;
|
||||
let session = self.sessions.get(&addr).context("Session not found")?;
|
||||
|
||||
let filename = format!(
|
||||
"temp/recorded_{}.wav",
|
||||
session.info.serial_number.replace(":", "")
|
||||
);
|
||||
|
||||
// 通知 Audio Manager 开始录音
|
||||
session
|
||||
.audio_manager
|
||||
.start_recording(config.clone(), filename)
|
||||
.await?;
|
||||
// 创建录音流,订阅音频总线
|
||||
let handle = RecorderStream::spawn(
|
||||
config.clone(),
|
||||
filename.clone(),
|
||||
self.audio_bus.subscribe(),
|
||||
Some(session.audio_addr),
|
||||
session.cancel.clone(),
|
||||
);
|
||||
|
||||
session.start_recording(handle, config.clone());
|
||||
|
||||
// 通知客户端开始发送音频
|
||||
session
|
||||
.conn
|
||||
.send(&ControlPacket::StartRecording { config })
|
||||
.await?;
|
||||
|
||||
// 发布事件
|
||||
session
|
||||
.send(&ControlPacket::ServerEvent(
|
||||
ServerEvent::AudioStatusChanged {
|
||||
is_recording: true,
|
||||
is_playing: session.is_playing(),
|
||||
},
|
||||
))
|
||||
.await?;
|
||||
|
||||
println!("[Server] Recording started for {} -> {}", addr, filename);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 停止录音
|
||||
pub async fn stop_record(&self, addr: SocketAddr) -> Result<()> {
|
||||
let session = self
|
||||
.sessions
|
||||
.get(&addr)
|
||||
.map(|r| r.value().clone())
|
||||
.context("Session not found")?;
|
||||
let session = self.sessions.get(&addr).context("Session not found")?;
|
||||
session.stop_recording();
|
||||
session.send(&ControlPacket::StopRecording).await?;
|
||||
|
||||
// 停止录音任务
|
||||
session.audio_manager.stop_recording().await?;
|
||||
session.conn.send(&ControlPacket::StopRecording).await?;
|
||||
// 发布事件
|
||||
session
|
||||
.send(&ControlPacket::ServerEvent(
|
||||
ServerEvent::AudioStatusChanged {
|
||||
is_recording: false,
|
||||
is_playing: session.is_playing(),
|
||||
},
|
||||
))
|
||||
.await?;
|
||||
|
||||
println!("[Server] Recording stopped for {}", addr);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 开始播放
|
||||
pub async fn start_play(&self, addr: SocketAddr, file_path: &str) -> Result<()> {
|
||||
let session = self
|
||||
.sessions
|
||||
.get(&addr)
|
||||
.map(|r| r.value().clone())
|
||||
.context("Session not found")?;
|
||||
let session = self.sessions.get(&addr).context("Session not found")?;
|
||||
|
||||
let reader = WavReader::open(file_path)?;
|
||||
let opus_rate = if reader.sample_rate > 24000 {
|
||||
@@ -314,31 +497,60 @@ impl Server {
|
||||
..AudioConfig::music_48k()
|
||||
};
|
||||
|
||||
// 通知客户端准备接收音频
|
||||
session
|
||||
.conn
|
||||
.send(&ControlPacket::StartPlayback {
|
||||
config: config.clone(),
|
||||
})
|
||||
.await?;
|
||||
|
||||
// 创建播放流
|
||||
let handle = FilePlaybackStream::spawn(
|
||||
config,
|
||||
reader,
|
||||
self.audio_bus.socket(),
|
||||
session.audio_addr,
|
||||
session.cancel.clone(),
|
||||
);
|
||||
|
||||
session.start_playback(handle);
|
||||
|
||||
// 发布事件
|
||||
session
|
||||
.audio_manager
|
||||
.start_playback(config, reader, self.audio.clone(), session.audio_addr)
|
||||
.send(&ControlPacket::ServerEvent(
|
||||
ServerEvent::AudioStatusChanged {
|
||||
is_recording: session.is_recording(),
|
||||
is_playing: true,
|
||||
},
|
||||
))
|
||||
.await?;
|
||||
|
||||
println!("[Server] Playback started for {} from {}", addr, file_path);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 停止播放
|
||||
pub async fn stop_play(&self, addr: SocketAddr) -> Result<()> {
|
||||
let session = self
|
||||
.sessions
|
||||
.get(&addr)
|
||||
.map(|r| r.value().clone())
|
||||
.context("Session not found")?;
|
||||
let session = self.sessions.get(&addr).context("Session not found")?;
|
||||
session.stop_playback();
|
||||
session.send(&ControlPacket::StopPlayback).await?;
|
||||
|
||||
// 停止播放任务
|
||||
session.audio_manager.stop_playback();
|
||||
session.conn.send(&ControlPacket::StopPlayback).await?;
|
||||
// 发布事件
|
||||
session
|
||||
.send(&ControlPacket::ServerEvent(
|
||||
ServerEvent::AudioStatusChanged {
|
||||
is_recording: session.is_recording(),
|
||||
is_playing: false,
|
||||
},
|
||||
))
|
||||
.await?;
|
||||
|
||||
println!("[Server] Playback stopped for {}", addr);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 关闭服务器
|
||||
pub fn shutdown(&self) {
|
||||
self.cancel.cancel();
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,326 @@
|
||||
//! # Session - 客户端会话管理
|
||||
//!
|
||||
//! 轻量级的会话结构,专注于:
|
||||
//! - TCP 控制连接
|
||||
//! - RPC 管理
|
||||
//! - 会话生命周期
|
||||
//!
|
||||
//! 音频流的实际处理由 AudioBus 和 Stream 模块负责。
|
||||
|
||||
use crate::audio::config::AudioConfig;
|
||||
use crate::net::command::{Command, CommandResult};
|
||||
use crate::net::network::Connection;
|
||||
use crate::net::protocol::{ClientInfo, ControlPacket};
|
||||
use crate::net::rpc::RpcManager;
|
||||
use anyhow::{Context, Result};
|
||||
use std::net::SocketAddr;
|
||||
use std::sync::Arc;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
use super::stream::StreamHandle;
|
||||
|
||||
/// 活动流追踪
|
||||
/// 用于追踪当前 session 的活动音频流
|
||||
pub struct ActiveStreams {
|
||||
/// 当前录音流句柄
|
||||
pub recorder: Option<StreamHandle>,
|
||||
/// 当前播放流句柄
|
||||
pub playback: Option<StreamHandle>,
|
||||
}
|
||||
|
||||
impl Default for ActiveStreams {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
recorder: None,
|
||||
playback: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl ActiveStreams {
|
||||
/// 停止所有活动流
|
||||
pub fn stop_all(&mut self) {
|
||||
if let Some(h) = self.recorder.take() {
|
||||
h.stop();
|
||||
}
|
||||
if let Some(h) = self.playback.take() {
|
||||
h.stop();
|
||||
}
|
||||
}
|
||||
|
||||
/// 开始录音(停止之前的录音)
|
||||
pub fn start_recording(&mut self, handle: StreamHandle) {
|
||||
if let Some(h) = self.recorder.take() {
|
||||
h.stop();
|
||||
}
|
||||
self.recorder = Some(handle);
|
||||
}
|
||||
|
||||
/// 停止录音
|
||||
pub fn stop_recording(&mut self) {
|
||||
if let Some(h) = self.recorder.take() {
|
||||
h.stop();
|
||||
}
|
||||
}
|
||||
|
||||
/// 开始播放(停止之前的播放)
|
||||
pub fn start_playback(&mut self, handle: StreamHandle) {
|
||||
if let Some(h) = self.playback.take() {
|
||||
h.stop();
|
||||
}
|
||||
self.playback = Some(handle);
|
||||
}
|
||||
|
||||
/// 停止播放
|
||||
pub fn stop_playback(&mut self) {
|
||||
if let Some(h) = self.playback.take() {
|
||||
h.stop();
|
||||
}
|
||||
}
|
||||
|
||||
/// 检查是否正在录音
|
||||
pub fn is_recording(&self) -> bool {
|
||||
self.recorder.is_some()
|
||||
}
|
||||
|
||||
/// 检查是否正在播放
|
||||
pub fn is_playing(&self) -> bool {
|
||||
self.playback.is_some()
|
||||
}
|
||||
}
|
||||
|
||||
/// 客户端会话
|
||||
pub struct Session {
|
||||
/// 客户端信息
|
||||
pub info: ClientInfo,
|
||||
|
||||
/// TCP 控制连接
|
||||
pub conn: Arc<Connection>,
|
||||
|
||||
/// RPC 管理器
|
||||
pub rpc: Arc<RpcManager>,
|
||||
|
||||
/// TCP 地址(用作会话 ID)
|
||||
pub tcp_addr: SocketAddr,
|
||||
|
||||
/// UDP 音频地址
|
||||
pub audio_addr: SocketAddr,
|
||||
|
||||
/// 会话取消令牌
|
||||
pub cancel: CancellationToken,
|
||||
|
||||
/// 活动流
|
||||
streams: parking_lot::Mutex<ActiveStreams>,
|
||||
|
||||
/// 当前录音配置(如果正在录音)
|
||||
recording_config: parking_lot::Mutex<Option<AudioConfig>>,
|
||||
|
||||
/// 会话创建时间
|
||||
created_at: std::time::Instant,
|
||||
}
|
||||
|
||||
impl Session {
|
||||
/// 创建新会话
|
||||
pub fn new(
|
||||
info: ClientInfo,
|
||||
conn: Arc<Connection>,
|
||||
tcp_addr: SocketAddr,
|
||||
audio_addr: SocketAddr,
|
||||
cancel: CancellationToken,
|
||||
) -> Self {
|
||||
Self {
|
||||
info,
|
||||
conn,
|
||||
rpc: Arc::new(RpcManager::new()),
|
||||
tcp_addr,
|
||||
audio_addr,
|
||||
cancel,
|
||||
streams: parking_lot::Mutex::new(ActiveStreams::default()),
|
||||
recording_config: parking_lot::Mutex::new(None),
|
||||
created_at: std::time::Instant::now(),
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取会话 ID(使用 TCP 地址)
|
||||
pub fn id(&self) -> SocketAddr {
|
||||
self.tcp_addr
|
||||
}
|
||||
|
||||
/// 检查会话是否仍然有效
|
||||
pub fn is_alive(&self) -> bool {
|
||||
!self.cancel.is_cancelled()
|
||||
}
|
||||
|
||||
/// 获取会话运行时间(秒)
|
||||
pub fn uptime_secs(&self) -> u64 {
|
||||
self.created_at.elapsed().as_secs()
|
||||
}
|
||||
|
||||
/// 发送控制包
|
||||
pub async fn send(&self, packet: &ControlPacket) -> Result<()> {
|
||||
self.conn.send(packet).await
|
||||
}
|
||||
|
||||
/// 接收控制包
|
||||
pub async fn recv(&self) -> Result<ControlPacket> {
|
||||
self.conn.recv().await
|
||||
}
|
||||
|
||||
/// 执行 RPC 调用(新版,使用 Command)
|
||||
pub async fn execute(&self, command: Command) -> Result<CommandResult> {
|
||||
let (id, rx) = self.rpc.register();
|
||||
self.conn
|
||||
.send(&ControlPacket::RpcRequest { id, command })
|
||||
.await?;
|
||||
rx.await.context("RPC channel closed")
|
||||
}
|
||||
|
||||
/// 处理 RPC 响应
|
||||
pub fn resolve_rpc(&self, id: u32, result: CommandResult) {
|
||||
self.rpc.resolve(id, result);
|
||||
}
|
||||
|
||||
/// 开始录音流
|
||||
pub fn start_recording(&self, handle: StreamHandle, config: AudioConfig) {
|
||||
self.streams.lock().start_recording(handle);
|
||||
*self.recording_config.lock() = Some(config);
|
||||
}
|
||||
|
||||
/// 停止录音流
|
||||
pub fn stop_recording(&self) {
|
||||
self.streams.lock().stop_recording();
|
||||
*self.recording_config.lock() = None;
|
||||
}
|
||||
|
||||
/// 开始播放流
|
||||
pub fn start_playback(&self, handle: StreamHandle) {
|
||||
self.streams.lock().start_playback(handle);
|
||||
}
|
||||
|
||||
/// 停止播放流
|
||||
pub fn stop_playback(&self) {
|
||||
self.streams.lock().stop_playback();
|
||||
}
|
||||
|
||||
/// 检查是否正在录音
|
||||
pub fn is_recording(&self) -> bool {
|
||||
self.streams.lock().is_recording()
|
||||
}
|
||||
|
||||
/// 检查是否正在播放
|
||||
pub fn is_playing(&self) -> bool {
|
||||
self.streams.lock().is_playing()
|
||||
}
|
||||
|
||||
/// 清理所有资源
|
||||
pub fn cleanup(&self) {
|
||||
self.cancel.cancel();
|
||||
self.streams.lock().stop_all();
|
||||
}
|
||||
|
||||
/// 获取当前录音配置
|
||||
pub fn recording_config(&self) -> Option<AudioConfig> {
|
||||
self.recording_config.lock().clone()
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for Session {
|
||||
fn drop(&mut self) {
|
||||
self.cleanup();
|
||||
}
|
||||
}
|
||||
|
||||
/// 会话管理器
|
||||
/// 负责管理所有客户端会话
|
||||
pub struct SessionManager {
|
||||
sessions: dashmap::DashMap<SocketAddr, Arc<Session>>,
|
||||
/// 从 UDP 地址到 TCP 地址的映射
|
||||
udp_to_tcp: dashmap::DashMap<SocketAddr, SocketAddr>,
|
||||
}
|
||||
|
||||
impl SessionManager {
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
sessions: dashmap::DashMap::new(),
|
||||
udp_to_tcp: dashmap::DashMap::new(),
|
||||
}
|
||||
}
|
||||
|
||||
/// 注册新会话
|
||||
pub fn register(&self, session: Arc<Session>) {
|
||||
let tcp_addr = session.tcp_addr;
|
||||
let audio_addr = session.audio_addr;
|
||||
|
||||
self.udp_to_tcp.insert(audio_addr, tcp_addr);
|
||||
self.sessions.insert(tcp_addr, session);
|
||||
|
||||
println!(
|
||||
"[SessionManager] Registered: {} (audio: {})",
|
||||
tcp_addr, audio_addr
|
||||
);
|
||||
}
|
||||
|
||||
/// 注销会话
|
||||
pub fn unregister(&self, tcp_addr: &SocketAddr) -> Option<Arc<Session>> {
|
||||
if let Some((_, session)) = self.sessions.remove(tcp_addr) {
|
||||
self.udp_to_tcp.remove(&session.audio_addr);
|
||||
session.cleanup();
|
||||
println!(
|
||||
"[SessionManager] Unregistered: {} ({})",
|
||||
tcp_addr, session.info.model
|
||||
);
|
||||
Some(session)
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
/// 通过 TCP 地址获取会话
|
||||
pub fn get(&self, tcp_addr: &SocketAddr) -> Option<Arc<Session>> {
|
||||
self.sessions.get(tcp_addr).map(|r| r.value().clone())
|
||||
}
|
||||
|
||||
/// 通过 UDP 地址获取会话
|
||||
pub fn get_by_udp(&self, udp_addr: &SocketAddr) -> Option<Arc<Session>> {
|
||||
self.udp_to_tcp
|
||||
.get(udp_addr)
|
||||
.and_then(|tcp_addr| self.get(tcp_addr.value()))
|
||||
}
|
||||
|
||||
/// 获取所有会话地址
|
||||
pub fn all_addrs(&self) -> Vec<SocketAddr> {
|
||||
self.sessions.iter().map(|r| *r.key()).collect()
|
||||
}
|
||||
|
||||
/// 获取所有会话
|
||||
pub fn all_sessions(&self) -> Vec<Arc<Session>> {
|
||||
self.sessions.iter().map(|r| r.value().clone()).collect()
|
||||
}
|
||||
|
||||
/// 获取会话数量
|
||||
pub fn count(&self) -> usize {
|
||||
self.sessions.len()
|
||||
}
|
||||
|
||||
/// 广播控制包到所有会话
|
||||
pub async fn broadcast(&self, packet: &ControlPacket) {
|
||||
for entry in self.sessions.iter() {
|
||||
let _ = entry.value().send(packet).await;
|
||||
}
|
||||
}
|
||||
|
||||
/// 广播控制包到所有会话(除了指定的)
|
||||
pub async fn broadcast_except(&self, packet: &ControlPacket, except: &SocketAddr) {
|
||||
for entry in self.sessions.iter() {
|
||||
if entry.key() != except {
|
||||
let _ = entry.value().send(packet).await;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for SessionManager {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,302 @@
|
||||
//! # AudioStream - 音频流抽象
|
||||
//!
|
||||
//! 提供统一的音频流处理接口,支持多种输入源和输出目标。
|
||||
//!
|
||||
//! ## 设计
|
||||
//!
|
||||
//! ```text
|
||||
//! ┌─────────────────────────────────────────────────────────────┐
|
||||
//! │ Stream Types │
|
||||
//! │ │
|
||||
//! │ ┌──────────────┐ ┌──────────────┐ ┌──────────────┐ │
|
||||
//! │ │ FileSource │ │ BusSource │ │ NetworkSink │ │
|
||||
//! │ │ (WAV Reader) │ │ (From Bus) │ │ (To Client) │ │
|
||||
//! │ └──────────────┘ └──────────────┘ └──────────────┘ │
|
||||
//! │ │
|
||||
//! │ ┌──────────────┐ ┌──────────────┐ │
|
||||
//! │ │ FileSink │ │ BusSink │ │
|
||||
//! │ │ (WAV Writer) │ │ (To Bus) │ │
|
||||
//! │ └──────────────┘ └──────────────┘ │
|
||||
//! └─────────────────────────────────────────────────────────────┘
|
||||
//! ```
|
||||
|
||||
use crate::audio::codec::OpusCodec;
|
||||
use crate::audio::config::AudioConfig;
|
||||
use crate::audio::wav::{WavReader, WavWriter};
|
||||
use crate::net::network::AudioSocket;
|
||||
use crate::net::protocol::AudioPacket;
|
||||
use std::net::SocketAddr;
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::broadcast;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
use super::audio_bus::AudioFrame;
|
||||
|
||||
/// 音频流任务句柄
|
||||
/// 用于控制正在运行的音频流任务
|
||||
pub struct StreamHandle {
|
||||
cancel: CancellationToken,
|
||||
}
|
||||
|
||||
impl StreamHandle {
|
||||
pub fn new(cancel: CancellationToken) -> Self {
|
||||
Self { cancel }
|
||||
}
|
||||
|
||||
/// 停止流
|
||||
pub fn stop(&self) {
|
||||
self.cancel.cancel();
|
||||
}
|
||||
|
||||
/// 检查是否已停止
|
||||
pub fn is_stopped(&self) -> bool {
|
||||
self.cancel.is_cancelled()
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for StreamHandle {
|
||||
fn drop(&mut self) {
|
||||
self.cancel.cancel();
|
||||
}
|
||||
}
|
||||
|
||||
/// 文件播放流 - 从 WAV 文件读取并发送到指定客户端
|
||||
pub struct FilePlaybackStream;
|
||||
|
||||
impl FilePlaybackStream {
|
||||
/// 启动文件播放流
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `config` - 音频配置(用于 Opus 编码)
|
||||
/// * `reader` - WAV 文件读取器
|
||||
/// * `socket` - UDP socket
|
||||
/// * `target` - 目标客户端地址
|
||||
/// * `parent_cancel` - 父级取消令牌(用于 session 级别取消)
|
||||
pub fn spawn(
|
||||
config: AudioConfig,
|
||||
reader: WavReader,
|
||||
socket: Arc<AudioSocket>,
|
||||
target: SocketAddr,
|
||||
parent_cancel: CancellationToken,
|
||||
) -> StreamHandle {
|
||||
let cancel = parent_cancel.child_token();
|
||||
let token = cancel.clone();
|
||||
|
||||
tokio::spawn(async move {
|
||||
if let Err(e) = Self::run(config, reader, socket, target, token).await {
|
||||
eprintln!("[FilePlayback] Error: {}", e);
|
||||
}
|
||||
});
|
||||
|
||||
StreamHandle::new(cancel)
|
||||
}
|
||||
|
||||
async fn run(
|
||||
config: AudioConfig,
|
||||
mut reader: WavReader,
|
||||
socket: Arc<AudioSocket>,
|
||||
target: SocketAddr,
|
||||
cancel: CancellationToken,
|
||||
) -> anyhow::Result<()> {
|
||||
let mut codec = OpusCodec::new(&config)?;
|
||||
let mut pcm = vec![0i16; config.frame_size];
|
||||
let mut opus_buf = vec![0u8; 4096];
|
||||
let frame_duration = std::time::Duration::from_millis(20);
|
||||
let mut interval = tokio::time::interval(frame_duration);
|
||||
|
||||
println!("[FilePlayback] Started -> {}", target);
|
||||
|
||||
loop {
|
||||
tokio::select! {
|
||||
_ = cancel.cancelled() => break,
|
||||
_ = interval.tick() => {
|
||||
match reader.read_samples(&mut pcm) {
|
||||
Ok(0) => {
|
||||
println!("[FilePlayback] EOF reached");
|
||||
break;
|
||||
}
|
||||
Ok(n) => {
|
||||
if let Ok(len) = codec.encode(&pcm[..n], &mut opus_buf) {
|
||||
let packet = AudioPacket {
|
||||
data: opus_buf[..len].to_vec(),
|
||||
};
|
||||
let _ = socket.send(&packet, target).await;
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
eprintln!("[FilePlayback] Read error: {}", e);
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
println!("[FilePlayback] Stopped");
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
/// 录音流 - 从总线订阅音频并写入 WAV 文件
|
||||
pub struct RecorderStream;
|
||||
|
||||
impl RecorderStream {
|
||||
/// 启动录音流
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `config` - 音频配置
|
||||
/// * `filename` - 输出文件路径
|
||||
/// * `bus_rx` - 音频总线接收器
|
||||
/// * `source_filter` - 仅录制来自此地址的音频(None 表示全部录制)
|
||||
/// * `parent_cancel` - 父级取消令牌
|
||||
pub fn spawn(
|
||||
config: AudioConfig,
|
||||
filename: String,
|
||||
bus_rx: broadcast::Receiver<AudioFrame>,
|
||||
source_filter: Option<SocketAddr>,
|
||||
parent_cancel: CancellationToken,
|
||||
) -> StreamHandle {
|
||||
let cancel = parent_cancel.child_token();
|
||||
let token = cancel.clone();
|
||||
|
||||
tokio::spawn(async move {
|
||||
if let Err(e) = Self::run(config, filename, bus_rx, source_filter, token).await {
|
||||
eprintln!("[Recorder] Error: {}", e);
|
||||
}
|
||||
});
|
||||
|
||||
StreamHandle::new(cancel)
|
||||
}
|
||||
|
||||
async fn run(
|
||||
config: AudioConfig,
|
||||
filename: String,
|
||||
mut bus_rx: broadcast::Receiver<AudioFrame>,
|
||||
source_filter: Option<SocketAddr>,
|
||||
cancel: CancellationToken,
|
||||
) -> anyhow::Result<()> {
|
||||
let mut writer = WavWriter::create(&filename, config.sample_rate, config.channels)?;
|
||||
let mut codec = OpusCodec::new(&config)?;
|
||||
let mut pcm = vec![0i16; config.frame_size];
|
||||
|
||||
println!(
|
||||
"[Recorder] Started -> {} (filter: {:?})",
|
||||
filename, source_filter
|
||||
);
|
||||
|
||||
loop {
|
||||
tokio::select! {
|
||||
_ = cancel.cancelled() => break,
|
||||
result = bus_rx.recv() => {
|
||||
match result {
|
||||
Ok(frame) => {
|
||||
// 应用源过滤
|
||||
if let Some(filter_addr) = &source_filter {
|
||||
if frame.source.as_ref() != Some(filter_addr) {
|
||||
continue;
|
||||
}
|
||||
}
|
||||
|
||||
// 解码并写入
|
||||
if let Ok(n) = codec.decode(&frame.packet.data, &mut pcm) {
|
||||
let _ = writer.write_samples(&pcm[..n]);
|
||||
}
|
||||
}
|
||||
Err(broadcast::error::RecvError::Lagged(n)) => {
|
||||
eprintln!("[Recorder] Lagged {} frames", n);
|
||||
}
|
||||
Err(broadcast::error::RecvError::Closed) => {
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 确保正确关闭文件
|
||||
writer.finalize()?;
|
||||
println!("[Recorder] Stopped, file saved: {}", filename);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
/// 音频转发流 - 从总线订阅并转发到指定客户端
|
||||
/// 用于实现"监听"功能或者服务端音频源推送
|
||||
pub struct ForwardStream;
|
||||
|
||||
impl ForwardStream {
|
||||
/// 启动转发流
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `socket` - UDP socket
|
||||
/// * `target` - 目标地址
|
||||
/// * `bus_rx` - 音频总线接收器
|
||||
/// * `source_filter` - 源过滤(可选)
|
||||
/// * `parent_cancel` - 父级取消令牌
|
||||
pub fn spawn(
|
||||
socket: Arc<AudioSocket>,
|
||||
target: SocketAddr,
|
||||
bus_rx: broadcast::Receiver<AudioFrame>,
|
||||
source_filter: Option<SocketAddr>,
|
||||
parent_cancel: CancellationToken,
|
||||
) -> StreamHandle {
|
||||
let cancel = parent_cancel.child_token();
|
||||
let token = cancel.clone();
|
||||
|
||||
tokio::spawn(async move {
|
||||
if let Err(e) = Self::run(socket, target, bus_rx, source_filter, token).await {
|
||||
eprintln!("[Forward] Error: {}", e);
|
||||
}
|
||||
});
|
||||
|
||||
StreamHandle::new(cancel)
|
||||
}
|
||||
|
||||
async fn run(
|
||||
socket: Arc<AudioSocket>,
|
||||
target: SocketAddr,
|
||||
mut bus_rx: broadcast::Receiver<AudioFrame>,
|
||||
source_filter: Option<SocketAddr>,
|
||||
cancel: CancellationToken,
|
||||
) -> anyhow::Result<()> {
|
||||
println!(
|
||||
"[Forward] Started -> {} (filter: {:?})",
|
||||
target, source_filter
|
||||
);
|
||||
|
||||
loop {
|
||||
tokio::select! {
|
||||
_ = cancel.cancelled() => break,
|
||||
result = bus_rx.recv() => {
|
||||
match result {
|
||||
Ok(frame) => {
|
||||
// 应用源过滤
|
||||
if let Some(filter_addr) = &source_filter {
|
||||
if frame.source.as_ref() != Some(filter_addr) {
|
||||
continue;
|
||||
}
|
||||
}
|
||||
|
||||
// 不转发给自己
|
||||
if frame.source.as_ref() == Some(&target) {
|
||||
continue;
|
||||
}
|
||||
|
||||
let _ = socket.send(&frame.packet, target).await;
|
||||
}
|
||||
Err(broadcast::error::RecvError::Lagged(n)) => {
|
||||
eprintln!("[Forward] Lagged {} frames", n);
|
||||
}
|
||||
Err(broadcast::error::RecvError::Closed) => {
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
println!("[Forward] Stopped");
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user