diff --git a/packages/client-v2/Cargo.lock b/packages/client-v2/Cargo.lock index 3f88ac8..e6cb3ef 100644 --- a/packages/client-v2/Cargo.lock +++ b/packages/client-v2/Cargo.lock @@ -108,6 +108,26 @@ version = "1.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "790eea4361631c5e7d22598ecd5723ff611904e3344ce8720784c93e3d83d40b" +[[package]] +name = "crossbeam-utils" +version = "0.8.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d0a5c400df2834b80a4c3327b3aad3a4c4cd4de0629063962b03235697506a28" + +[[package]] +name = "dashmap" +version = "6.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5041cc499144891f3790297212f32a74fb938e5136a14943f338ef9e0ae276cf" +dependencies = [ + "cfg-if", + "crossbeam-utils", + "hashbrown 0.14.5", + "lock_api", + "once_cell", + "parking_lot_core", +] + [[package]] name = "embedded-io" version = "0.4.0" @@ -127,7 +147,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" dependencies = [ "libc", - "windows-sys 0.60.2", + "windows-sys 0.61.2", ] [[package]] @@ -136,6 +156,55 @@ version = "0.1.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "645cbb3a84e60b7531617d5ae4e57f7e27308f6445f5abf653209ea76dec8dff" +[[package]] +name = "futures-core" +version = "0.3.31" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "05f29059c0c2090612e8d742178b0580d2dc940c837851ad723096f87af6663e" + +[[package]] +name = "futures-io" +version = "0.3.31" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9e5c1b78ca4aae1ac06c48a526a655760685149f0d465d21f37abfe57ce075c6" + +[[package]] +name = "futures-macro" +version = "0.3.31" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "162ee34ebcb7c64a8abebc059ce0fee27c2262618d7b60ed8faf72fef13c3650" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "futures-sink" +version = "0.3.31" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e575fab7d1e0dcb8d0c7bcf9a63ee213816ab51902e6d244a95819acacf1d4f7" + +[[package]] +name = "futures-task" +version = "0.3.31" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f90f7dce0722e95104fcb095585910c0977252f286e354b5e3bd38902cd99988" + +[[package]] +name = "futures-util" +version = "0.3.31" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9fa08315bb612088cc391249efdc3bc77536f16c91f6cf495e6fbe85b20a4a81" +dependencies = [ + "futures-core", + "futures-macro", + "futures-task", + "pin-project-lite", + "pin-utils", + "slab", +] + [[package]] name = "hash32" version = "0.2.1" @@ -145,6 +214,18 @@ dependencies = [ "byteorder", ] +[[package]] +name = "hashbrown" +version = "0.14.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e5274423e17b7c9fc20b6e7e208532f9b19825d82dfd615708b70edd83df41f1" + +[[package]] +name = "hashbrown" +version = "0.15.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9229cfe53dfd69f0609a49f65461bd93001ea1ef889cd5529dd176593f5338a1" + [[package]] name = "heapless" version = "0.7.17" @@ -191,6 +272,12 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "once_cell" +version = "1.21.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "42f5e15c9953c5e4ccceeb2e7382a716482c34515315f7b03532b8b4e8393d2d" + [[package]] name = "opus" version = "0.3.0" @@ -230,6 +317,12 @@ version = "0.2.16" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3b3cff922bd51709b605d9ead9aa71031d81447142d828eb4a6eba76fe619f9b" +[[package]] +name = "pin-utils" +version = "0.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b870d8c151b6f2fb93e84a13146138f05d02ed11c7e7c54f8826aaaf7c9f184" + [[package]] name = "pkg-config" version = "0.3.32" @@ -343,6 +436,12 @@ dependencies = [ "libc", ] +[[package]] +name = "slab" +version = "0.4.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7a2ae44ef20feb57a68b23d846850f861394c2e02dc425a50098ae8c90267589" + [[package]] name = "smallvec" version = "1.15.1" @@ -433,6 +532,23 @@ dependencies = [ "syn", ] +[[package]] +name = "tokio-util" +version = "0.7.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2efa149fe76073d6e8fd97ef4f4eca7b67f599660115591483572e406e165594" +dependencies = [ + "bytes", + "futures-core", + "futures-io", + "futures-sink", + "futures-util", + "hashbrown 0.15.5", + "pin-project-lite", + "slab", + "tokio", +] + [[package]] name = "unicode-ident" version = "1.0.22" @@ -540,9 +656,11 @@ version = "0.1.0" dependencies = [ "alsa", "anyhow", + "dashmap", "opus", "parking_lot", "postcard", "serde", "tokio", + "tokio-util", ] diff --git a/packages/client-v2/Cargo.toml b/packages/client-v2/Cargo.toml index ab96699..0b8b87b 100644 --- a/packages/client-v2/Cargo.toml +++ b/packages/client-v2/Cargo.toml @@ -19,9 +19,11 @@ app = [] opus = "0.3" anyhow = "1.0" tokio = { version = "1.48", features = ["full"] } +tokio-util = { version = "0.7", features = ["full"] } serde = { version = "1.0", features = ["derive"] } postcard = { version = "1.0", features = ["alloc", "use-std"] } parking_lot = "0.12.5" +dashmap = "6.1.0" [target.'cfg(target_os = "linux")'.dependencies] alsa = "0.11" diff --git a/packages/client-v2/src/app/client/mod.rs b/packages/client-v2/src/app/client/mod.rs index 0bc401b..8c21e78 100644 --- a/packages/client-v2/src/app/client/mod.rs +++ b/packages/client-v2/src/app/client/mod.rs @@ -1,218 +1,227 @@ #![cfg(target_os = "linux")] use crate::audio::codec::OpusCodec; +use crate::audio::config::AudioConfig; use crate::audio::player::AudioPlayer; use crate::audio::recorder::AudioRecorder; use crate::net::discovery::Discovery; use crate::net::network::{AudioSocket, Connection}; -use crate::net::protocol::{AudioPacket, ControlPacket, DeviceInfo, RpcResult}; +use crate::net::protocol::{AudioPacket, ClientInfo, ControlPacket, RpcResult}; use crate::net::rpc::RpcManager; -use anyhow::Result; +use anyhow::{Result, anyhow}; use std::net::SocketAddr; use std::sync::Arc; -use tokio::sync::{Mutex, broadcast}; +use tokio::sync::{RwLock, mpsc}; +use tokio_util::sync::CancellationToken; + +/// 内部连接上下文,包含了音频流所需的全部信息 +struct ActiveSession { + conn: Arc, + audio_socket: Arc, + server_audio_addr: SocketAddr, + session_cancel: CancellationToken, // 控制整个 Session 的生命周期 + record_cancel: RwLock>, + play_cancel: RwLock>, +} pub struct Client { - info: DeviceInfo, - conn: Mutex>>, + session: RwLock>>, rpc: Arc, - server_audio_addr: Mutex>, } impl Client { pub fn new() -> Self { Self { - info: DeviceInfo::current(), - conn: Mutex::new(None), + session: RwLock::new(None), rpc: Arc::new(RpcManager::new()), - server_audio_addr: Mutex::new(None), } } pub async fn run(self: Arc) -> Result<()> { - let (ip, tcp_port, udp_port) = Discovery::listen().await?; - let addr = SocketAddr::new(ip, tcp_port); - let audio_addr = SocketAddr::new(ip, udp_port); - println!("Found server at {}, audio at {}", addr, audio_addr); - *self.server_audio_addr.lock().await = Some(audio_addr); + loop { + let (ip, tcp_port) = match Discovery::listen().await { + Ok(res) => res, + Err(e) => { + eprintln!("Discovery listen error: {}", e); + tokio::time::sleep(std::time::Duration::from_secs(2)).await; + continue; + } + }; - println!("Connecting to TCP server at {}...", addr); - let stream = tokio::net::TcpStream::connect(addr).await?; + let addr = SocketAddr::new(ip, tcp_port); + println!("Found server at {}", addr); + + if let Ok(stream) = tokio::net::TcpStream::connect(addr).await { + println!("Connected to TCP server at {}", addr); + if let Err(e) = self.clone().handle_session(stream, addr).await { + eprintln!("Session error: {:?}", e); + } + } else { + eprintln!("Failed to connect to {}", addr); + } + + self.cleanup().await; + println!( + "Connection to {} closed, searching for server again...", + addr + ); + tokio::time::sleep(std::time::Duration::from_secs(1)).await; + } + } + + async fn cleanup(&self) { + let mut session_guard = self.session.write().await; + if let Some(s) = session_guard.take() { + s.session_cancel.cancel(); + } + } + + async fn handle_session( + self: Arc, + stream: tokio::net::TcpStream, + addr: SocketAddr, + ) -> Result<()> { let conn = Arc::new(Connection::new(stream)?); - println!("TCP connected, sending identification..."); + let audio_socket = Arc::new(AudioSocket::bind().await?); - let audio = Arc::new(AudioSocket::bind().await?); - conn.send(&ControlPacket::ClientIdentify { - info: self.info.clone(), - udp_port: audio.port(), + // --- 握手 (Handshake) --- + let version = env!("CARGO_PKG_VERSION").to_string(); + let server_auth = + std::env::var("XIAO_SERVER_AUTH").unwrap_or_else(|_| "open-xiaoai".to_string()); + let client_auth = + std::env::var("XIAO_CLIENT_AUTH").unwrap_or_else(|_| "open-xiaoai".to_string()); + + conn.send(&ControlPacket::ClientHello { + auth: server_auth, + version: version.clone(), + udp_port: audio_socket.port(), + info: ClientInfo { + model: "Open-XiaoAi-V2".to_string(), + serial_number: "00:00:00:00:00:00".to_string(), + }, }) .await?; - match conn.recv().await? { - ControlPacket::IdentifyOk => println!("Connected to server"), - p => return Err(anyhow::anyhow!("Handshake failed: {:?}", p)), - } - *self.conn.lock().await = Some(conn.clone()); - let (stop_tx, _) = broadcast::channel(1); - - loop { - let packet = conn.recv().await?; - let this = self.clone(); - let audio = audio.clone(); - let stop_tx = stop_tx.clone(); - let audio_addr = self.server_audio_addr.lock().await.unwrap(); - tokio::spawn(async move { - if let Err(e) = this.handle_packet(packet, audio, stop_tx, audio_addr).await { - eprintln!("Handle packet error: {}", e); + let server_udp_port = match conn.recv().await? { + ControlPacket::ServerHello { + version: v, + udp_port, + auth, + } => { + if v != version { + return Err(anyhow!("Server version mismatch: {} != {}", v, version)); } - }); + if auth != client_auth { + return Err(anyhow!("Invalid server auth")); + } + udp_port + } + _ => return Err(anyhow!("Handshake failed")), + }; + + let server_audio_addr = SocketAddr::new(addr.ip(), server_udp_port); + println!( + "Handshake successful with {}, audio at {}", + addr, server_audio_addr + ); + + // --- 初始化 Session --- + let session = Arc::new(ActiveSession { + conn: conn.clone(), + audio_socket, + server_audio_addr, + session_cancel: CancellationToken::new(), + record_cancel: RwLock::new(None), + play_cancel: RwLock::new(None), + }); + *self.session.write().await = Some(session.clone()); + + // 使用 mpsc 队列来缓冲指令,确保顺序执行且不阻塞接收循环 + let (cmd_tx, mut cmd_rx) = mpsc::channel::(64); + + // 任务 1: 心跳 + let hb_session = session.clone(); + tokio::spawn(async move { + let mut interval = tokio::time::interval(std::time::Duration::from_secs(10)); + loop { + tokio::select! { + _ = hb_session.session_cancel.cancelled() => break, + _ = interval.tick() => { + if hb_session.conn.send(&ControlPacket::Ping).await.is_err() { break; } + } + } + } + }); + + // 任务 2: 命令处理器 (串行处理所有指令) + let proc_self = self.clone(); + let proc_session = session.clone(); + tokio::spawn(async move { + while let Some(packet) = cmd_rx.recv().await { + if let Err(e) = proc_self.process_packet(packet, &proc_session).await { + eprintln!("Process error: {}", e); + } + } + }); + + // 任务 3: 接收循环 (高优先级,只读包并分发) + loop { + tokio::select! { + _ = session.session_cancel.cancelled() => break, + res = tokio::time::timeout(std::time::Duration::from_secs(60), conn.recv()) => { + match res { + Ok(Ok(packet)) => { + if cmd_tx.send(packet).await.is_err() { break; } + } + Ok(Err(e)) => { + return Err(anyhow!("Connection receive error: {}", e)); + } + Err(_) => return Err(anyhow!("Connection timeout")), + } + } + } } + + Ok(()) } - pub async fn call(&self, method: &str, args: Vec) -> Result { - let (id, rx) = self.rpc.register(); - if let Some(conn) = self.conn.lock().await.as_ref() { - conn.send(&ControlPacket::RpcRequest { - id, - method: method.to_string(), - args, - }) - .await?; - Ok(rx.await?) - } else { - Err(anyhow::anyhow!("Not connected")) - } - } - - async fn handle_packet( + async fn process_packet( &self, packet: ControlPacket, - audio: Arc, - stop_tx: broadcast::Sender<()>, - server_addr: SocketAddr, + session: &Arc, ) -> Result<()> { match packet { + ControlPacket::Ping => { + session.conn.send(&ControlPacket::Pong).await?; + } + ControlPacket::Pong => {} + ControlPacket::RpcResponse { id, result } => { + self.rpc.resolve(id, result); + } ControlPacket::RpcRequest { id, method, args } => { let result = self.handle_rpc(&method, args).await; - if let Some(conn) = self.conn.lock().await.as_ref() { - conn.send(&ControlPacket::RpcResponse { id, result }) - .await?; - } + session + .conn + .send(&ControlPacket::RpcResponse { id, result }) + .await?; } - ControlPacket::RpcResponse { id, result } => self.rpc.resolve(id, result), ControlPacket::StartRecording { config } => { - let mut stop_rx = stop_tx.subscribe(); - tokio::spawn(async move { - let (pcm_tx, mut pcm_rx) = tokio::sync::mpsc::channel::>(20); - - // 录音线程:使用 std::thread 处理阻塞的 ALSA 调用 - let config_clone = config.clone(); - std::thread::spawn(move || { - let recorder = match AudioRecorder::new(&config_clone) { - Ok(r) => r, - Err(e) => { - eprintln!("Failed to start recorder: {}", e); - return; - } - }; - loop { - let mut pcm = vec![0i16; config_clone.frame_size]; - match recorder.read(&mut pcm) { - Ok(n) => { - if pcm_tx.blocking_send(pcm[..n].to_vec()).is_err() { - break; // Receiver dropped, stop recording - } - } - Err(e) => { - eprintln!("Recorder read error: {}", e); - break; - } - } - } - }); - - let mut codec = match OpusCodec::new(&config) { - Ok(c) => c, - Err(e) => { - eprintln!("Failed to init opus codec: {}", e); - return; - } - }; - - println!("Recording started..."); - loop { - tokio::select! { - _ = stop_rx.recv() => break, - Some(pcm_data) = pcm_rx.recv() => { - let mut opus = vec![0u8; 4096]; - if let Ok(len) = codec.encode(&pcm_data, &mut opus) { - let _ = audio.send(&AudioPacket { data: opus[..len].to_vec() }, server_addr).await; - } - } - } - } - println!("Recording stopped."); - }); + self.stop_recorder(session).await; // 开启前先停止旧的,防止资源冲突 + let token = session.session_cancel.child_token(); + *session.record_cancel.write().await = Some(token.clone()); + self.spawn_recorder(session.clone(), config, token); } ControlPacket::StartPlayback { config } => { - let mut stop_rx = stop_tx.subscribe(); - tokio::spawn(async move { - let (pcm_tx, mut pcm_rx) = tokio::sync::mpsc::channel::>(20); - - // 播放线程:使用 std::thread 处理阻塞的 ALSA 调用 - let config_clone = config.clone(); - std::thread::spawn(move || { - let player = match AudioPlayer::new(&config_clone) { - Ok(p) => p, - Err(e) => { - eprintln!("Failed to start player: {}", e); - return; - } - }; - while let Some(pcm_data) = pcm_rx.blocking_recv() { - let _ = player.write(&pcm_data); - } - }); - - let mut codec = match OpusCodec::new(&config) { - Ok(c) => c, - Err(e) => { - eprintln!("Failed to init opus codec: {}", e); - return; - } - }; - - let mut buf = vec![0u8; 4096]; - println!("Playback started with jitter buffer..."); - loop { - tokio::select! { - _ = stop_rx.recv() => break, - res = audio.recv(&mut buf) => { - match res { - Ok((packet, _)) => { - let mut pcm = vec![0i16; config.frame_size]; - if let Ok(n) = codec.decode(&packet.data, &mut pcm) { - let _ = pcm_tx.send(pcm[..n].to_vec()).await; - } - } - Err(e) => { - eprintln!("Audio recv error: {}", e); - break; - } - } - } - } - } - println!("Playback stopped."); - }); + self.stop_player(session).await; + let token = session.session_cancel.child_token(); + *session.play_cancel.write().await = Some(token.clone()); + self.spawn_player(session.clone(), config, token); } - ControlPacket::StopRecording | ControlPacket::StopPlayback => { - let _ = stop_tx.send(()); + ControlPacket::StopRecording => { + self.stop_recorder(session).await; } - ControlPacket::Ping => { - if let Some(conn) = self.conn.lock().await.as_ref() { - conn.send(&ControlPacket::Pong).await?; - } + ControlPacket::StopPlayback => { + self.stop_player(session).await; } _ => {} } @@ -246,4 +255,136 @@ impl Client { }, } } + + async fn stop_recorder(&self, session: &ActiveSession) { + let mut cancel_guard = session.record_cancel.write().await; + if let Some(token) = cancel_guard.take() { + token.cancel(); + } + } + + async fn stop_player(&self, session: &ActiveSession) { + let mut cancel_guard = session.play_cancel.write().await; + if let Some(token) = cancel_guard.take() { + token.cancel(); + } + } + + fn spawn_recorder( + &self, + session: Arc, + config: AudioConfig, + token: CancellationToken, + ) { + tokio::spawn(async move { + let (pcm_tx, mut pcm_rx) = mpsc::channel::>(32); + let conf = config.clone(); + + // 录音线程 (ALSA 阻塞) + std::thread::spawn(move || { + let recorder = match AudioRecorder::new(&conf) { + Ok(r) => r, + Err(e) => { + eprintln!("Failed to start recorder: {}", e); + return; + } + }; + let mut buf = vec![0i16; conf.frame_size]; + while let Ok(n) = recorder.read(&mut buf) { + if pcm_tx.blocking_send(buf[..n].to_vec()).is_err() { + break; + } + } + }); + + let mut codec = match OpusCodec::new(&config) { + Ok(c) => c, + Err(e) => { + eprintln!("Failed to init opus codec: {}", e); + return; + } + }; + println!("Recording started..."); + loop { + tokio::select! { + _ = token.cancelled() => break, + Some(pcm) = pcm_rx.recv() => { + let mut out = vec![0u8; 4096]; + if let Ok(len) = codec.encode(&pcm, &mut out) { + let _ = session.audio_socket.send(&AudioPacket { data: out[..len].to_vec() }, session.server_audio_addr).await; + } + } + } + } + println!("Recording stopped."); + }); + } + + fn spawn_player( + &self, + session: Arc, + config: AudioConfig, + token: CancellationToken, + ) { + tokio::spawn(async move { + let (pcm_tx, mut pcm_rx) = mpsc::channel::>(32); + let conf = config.clone(); + + // 播放线程 (ALSA 阻塞) + std::thread::spawn(move || { + let player = match AudioPlayer::new(&conf) { + Ok(p) => p, + Err(e) => { + eprintln!("Failed to start player: {}", e); + return; + } + }; + while let Some(pcm) = pcm_rx.blocking_recv() { + let _ = player.write(&pcm); + } + }); + + let mut codec = match OpusCodec::new(&config) { + Ok(c) => c, + Err(e) => { + eprintln!("Failed to init opus codec: {}", e); + return; + } + }; + let mut udp_buf = vec![0u8; 4096]; + println!("Playback started..."); + loop { + tokio::select! { + _ = token.cancelled() => break, + res = session.audio_socket.recv(&mut udp_buf) => { + if let Ok((packet, _)) = res { + let mut pcm = vec![0i16; config.frame_size]; + if let Ok(n) = codec.decode(&packet.data, &mut pcm) { + let _ = pcm_tx.send(pcm[..n].to_vec()).await; + } + } + } + } + } + println!("Playback stopped."); + }); + } + + pub async fn call(&self, method: &str, args: Vec) -> Result { + let (id, rx) = self.rpc.register(); + let session_guard = self.session.read().await; + if let Some(session) = session_guard.as_ref() { + session + .conn + .send(&ControlPacket::RpcRequest { + id, + method: method.to_string(), + args, + }) + .await?; + Ok(rx.await?) + } else { + Err(anyhow!("Not connected")) + } + } } diff --git a/packages/client-v2/src/app/mod.rs b/packages/client-v2/src/app/mod.rs index c07f47e..4574649 100644 --- a/packages/client-v2/src/app/mod.rs +++ b/packages/client-v2/src/app/mod.rs @@ -1,2 +1,3 @@ +#[cfg(target_os = "linux")] pub mod client; pub mod server; diff --git a/packages/client-v2/src/app/server/mod.rs b/packages/client-v2/src/app/server/mod.rs index 60162c9..0e61a8e 100644 --- a/packages/client-v2/src/app/server/mod.rs +++ b/packages/client-v2/src/app/server/mod.rs @@ -3,30 +3,50 @@ use crate::audio::config::AudioConfig; use crate::audio::wav::{WavReader, WavWriter}; use crate::net::discovery::Discovery; use crate::net::network::{AudioSocket, Connection}; -use crate::net::protocol::{AudioPacket, ControlPacket, DeviceInfo, RpcResult}; +use crate::net::protocol::{AudioPacket, ClientInfo, ControlPacket, RpcResult}; use crate::net::rpc::RpcManager; -use anyhow::{Context, Result}; -use std::collections::HashMap; +use anyhow::{Context, Result, anyhow}; +use dashmap::DashMap; +use parking_lot::Mutex; use std::net::SocketAddr; use std::sync::Arc; -use tokio::sync::Mutex; +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 Session { - pub info: DeviceInfo, + pub info: ClientInfo, pub conn: Arc, pub rpc: Arc, + pub tcp_addr: SocketAddr, pub audio_addr: SocketAddr, + pub session_cancel: CancellationToken, + pub record_cancel: Mutex>, + pub play_cancel: Mutex>, + pub audio_tx: mpsc::Sender, + pub recorder_tx: mpsc::Sender, + pub tracker: TaskTracker, } pub struct Server { - sessions: Arc>>>, + sessions: DashMap>, + udp_to_tcp: DashMap, audio: Arc, } impl Server { pub async fn new() -> Result { Ok(Self { - sessions: Arc::new(Mutex::new(HashMap::new())), + sessions: DashMap::new(), + udp_to_tcp: DashMap::new(), audio: Arc::new(AudioSocket::bind().await?), }) } @@ -35,65 +55,228 @@ impl Server { let listener = tokio::net::TcpListener::bind(format!("0.0.0.0:{}", port)).await?; let addr = listener.local_addr()?; println!("Server listening on TCP: {}", addr); - Discovery::broadcast(port, self.audio.port()).await?; + Discovery::broadcast(port).await?; + + // Audio Dispatcher + let this = self.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_tx.try_send(packet); + } + } + } + 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.handle_connection(stream, addr).await { + if let Err(e) = server.clone().handle_connection(stream, addr).await { eprintln!("Session {} error: {}", addr, e); } - server.sessions.lock().await.remove(&addr); + server.remove_session(&addr).await; }); } } + 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_connection( - &self, + self: Arc, stream: tokio::net::TcpStream, addr: SocketAddr, ) -> Result<()> { println!("New TCP connection from {}", addr); let conn = Arc::new(Connection::new(stream)?); - let (info, client_udp_port) = match conn.recv().await? { - ControlPacket::ClientIdentify { info, udp_port } => (info, udp_port), + // --- 握手 (Handshake) --- + let version = env!("CARGO_PKG_VERSION").to_string(); + let server_auth = + std::env::var("XIAO_SERVER_AUTH").unwrap_or_else(|_| "open-xiaoai".to_string()); + let client_auth = + std::env::var("XIAO_CLIENT_AUTH").unwrap_or_else(|_| "open-xiaoai".to_string()); + + let (info, client_audio_port) = match conn.recv().await? { + ControlPacket::ClientHello { + auth, + version: v, + udp_port, + info, + } => { + if v != version { + return Err(anyhow!("Client version mismatch: {} != {}", v, version)); + } + if auth != server_auth { + return Err(anyhow!("Invalid client auth")); + } + (info, udp_port) + } p => { - println!("Expected Identify from {}, got {:?}", addr, p); - return Err(anyhow::anyhow!("Expected Identify, got {:?}", p)); + return Err(anyhow!("Handshake failed: unexpected packet {:?}", p)); } }; - let audio_addr = SocketAddr::new(addr.ip(), client_udp_port); - println!( - "Client identified: {} ({}) version {}, audio at {}", - info.model, addr, info.version, audio_addr - ); - conn.send(&ControlPacket::IdentifyOk).await?; + conn.send(&ControlPacket::ServerHello { + auth: client_auth, + version: env!("CARGO_PKG_VERSION").to_string(), + udp_port: self.audio.port(), + }) + .await?; + let audio_addr = SocketAddr::new(addr.ip(), client_audio_port); + println!( + "Client identified: {} ({}), audio at {}", + info.model, info.serial_number, audio_addr + ); + + let (audio_tx, audio_rx) = mpsc::channel(1024); + let (recorder_tx, recorder_rx) = mpsc::channel(64); + let tracker = TaskTracker::new(); let session = Arc::new(Session { info, conn: conn.clone(), rpc: Arc::new(RpcManager::new()), + tcp_addr: addr, audio_addr, + session_cancel: CancellationToken::new(), + record_cancel: Mutex::new(None), + play_cancel: Mutex::new(None), + audio_tx, + recorder_tx, + tracker: tracker.clone(), }); - self.sessions.lock().await.insert(addr, session.clone()); + self.sessions.insert(addr, session.clone()); + self.udp_to_tcp.insert(audio_addr, addr); - loop { - let packet = conn.recv().await?; - let session = session.clone(); - tokio::spawn(async move { - if let Err(e) = handle_packet(session, packet).await { - eprintln!("Handle packet error: {}", e); + // --- Audio Processor Task --- + // 优化 Audio Receiver 的 OwnerShip,并使用 Tracker 跟踪 + let processor_session = session.clone(); + tracker.spawn(async move { + let mut active_recorder: Option<(WavWriter, OpusCodec, usize)> = None; + let mut audio_rx = audio_rx; + let mut recorder_rx = recorder_rx; + loop { + tokio::select! { + _ = processor_session.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(); + } + }); + + // --- Main Connection Loop --- + // 统一心跳与超时处理,节省资源 + let mut heartbeat = tokio::time::interval(std::time::Duration::from_secs(30)); + loop { + tokio::select! { + _ = session.session_cancel.cancelled() => break, + _ = heartbeat.tick() => { + if conn.send(&ControlPacket::Ping).await.is_err() { break; } + } + res = tokio::time::timeout(std::time::Duration::from_secs(60), conn.recv()) => { + match res { + Ok(Ok(packet)) => { + if let Err(e) = self.process_packet(session.clone(), packet).await { + eprintln!("Process packet error: {}", e); + } + } + Ok(Err(e)) => { + eprintln!("Session {} connection error: {}", addr, e); + break; + } + Err(_) => { + eprintln!("Session {} timeout", addr); + break; + } + } + } + } } + + Ok(()) + } + + async fn process_packet(&self, session: Arc, packet: ControlPacket) -> Result<()> { + match packet { + ControlPacket::Ping => { + session.conn.send(&ControlPacket::Pong).await?; + } + ControlPacket::Pong => {} + ControlPacket::RpcResponse { id, result } => { + session.rpc.resolve(id, result); + } + ControlPacket::RpcRequest { id, method, args } => { + let result = self.handle_rpc(&session, &method, args).await; + session + .conn + .send(&ControlPacket::RpcResponse { id, result }) + .await?; + } + _ => {} + } + Ok(()) } pub async fn get_clients(&self) -> Vec { - self.sessions.lock().await.keys().cloned().collect() + self.sessions.iter().map(|r| *r.key()).collect() } pub async fn call( @@ -104,10 +287,8 @@ impl Server { ) -> Result { let session = self .sessions - .lock() - .await .get(&addr) - .cloned() + .map(|r| r.value().clone()) .context("Session not found")?; let (id, rx) = session.rpc.register(); session @@ -121,59 +302,88 @@ impl Server { Ok(rx.await?) } + async fn handle_rpc( + &self, + _session: &Arc, + method: &str, + args: Vec, + ) -> RpcResult { + match method { + "hello" => RpcResult { + stdout: format!("Hello from server! Args: {:?}", args), + ..Default::default() + }, + _ => RpcResult { + stderr: format!("Unknown method: {}", method), + code: -1, + ..Default::default() + }, + } + } + pub async fn start_record(&self, addr: SocketAddr, config: AudioConfig) -> Result<()> { let session = self .sessions - .lock() - .await .get(&addr) - .cloned() + .map(|r| r.value().clone()) .context("Session not found")?; + + // 仅停止之前的录音任务 + { + let mut guard = session.record_cancel.lock(); + if let Some(token) = guard.take() { + token.cancel(); + } + *guard = Some(session.session_cancel.child_token()); + } + + let filename = format!( + "temp/recorded_{}.wav", + session.info.serial_number.replace(":", "") + ); + + // 通知 Audio Processor 开始录音 session - .conn - .send(&ControlPacket::StartRecording { + .recorder_tx + .send(RecorderCommand::Start { config: config.clone(), + filename, }) .await?; - let audio = self.audio.clone(); - let target_addr = session.audio_addr; - let filename = format!("temp/recorded_{}.wav", session.info.mac.replace(":", "")); - tokio::spawn(async move { - let mut writer = - WavWriter::create(&filename, config.sample_rate, config.channels).unwrap(); - let mut codec = OpusCodec::new(&config).unwrap(); - let mut pcm = vec![0i16; config.frame_size]; - let mut buf = vec![0u8; 4096]; - println!( - "Recording started for {}, saving to {}", - target_addr, filename - ); - for _ in 0..500 { - // Record ~10s - if let Ok((packet, src_addr)) = audio.recv(&mut buf).await { - if src_addr == target_addr { - if let Ok(n) = codec.decode(&packet.data, &mut pcm) { - writer.write_samples(&pcm[..n]).unwrap(); - } - } - } - } - writer.finalize().unwrap(); - println!("Recording saved to {}", filename); - }); + + session + .conn + .send(&ControlPacket::StartRecording { config }) + .await?; + Ok(()) } - pub async fn start_play(&self, addr: SocketAddr) -> Result<()> { + pub async fn stop_record(&self, addr: SocketAddr) -> Result<()> { let session = self .sessions - .lock() - .await .get(&addr) - .cloned() + .map(|r| r.value().clone()) .context("Session not found")?; - let reader = WavReader::open("temp/test.wav")?; + // 停止录音任务 + if let Some(token) = session.record_cancel.lock().take() { + token.cancel(); + } + + let _ = session.recorder_tx.send(RecorderCommand::Stop).await; + session.conn.send(&ControlPacket::StopRecording).await?; + 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 reader = WavReader::open(file_path)?; let opus_rate = if reader.sample_rate > 24000 { 48000 } else { @@ -187,66 +397,89 @@ impl Server { ..AudioConfig::music_48k() }; + // 仅停止之前的播放任务 + let token = { + let mut guard = session.play_cancel.lock(); + if let Some(token) = guard.take() { + token.cancel(); + } + let token = session.session_cancel.child_token(); + *guard = Some(token.clone()); + token + }; + session .conn .send(&ControlPacket::StartPlayback { config: config.clone(), }) .await?; - let audio = self.audio.clone(); + + let audio_socket = self.audio.clone(); let target_addr = session.audio_addr; - tokio::spawn(async move { + let playback_session = session.clone(); + + session.tracker.spawn(async move { let mut reader = reader; - let mut codec = OpusCodec::new(&config).unwrap(); + 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]; - while let Ok(n) = reader.read_samples(&mut pcm) { - if n == 0 { - break; + let mut interval = tokio::time::interval(std::time::Duration::from_millis(20)); + + println!("Playback started for {}", playback_session.tcp_addr); + + loop { + tokio::select! { + _ = token.cancelled() => break, + _ = playback_session.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; + } + } + } } - if let Ok(len) = codec.encode(&pcm[..n], &mut opus) { - let _ = audio - .send( - &AudioPacket { - data: opus[..len].to_vec(), - }, - target_addr, - ) - .await; - } - tokio::time::sleep(std::time::Duration::from_millis(20)).await; } + println!("Playback finished for {}", playback_session.tcp_addr); }); + + 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")?; + + // 停止播放任务 + if let Some(token) = session.play_cancel.lock().take() { + token.cancel(); + } + + session.conn.send(&ControlPacket::StopPlayback).await?; Ok(()) } } - -async fn handle_packet(session: Arc, packet: ControlPacket) -> Result<()> { - match packet { - ControlPacket::RpcRequest { id, method, args } => { - let result = handle_rpc(&session, &method, args).await; - session - .conn - .send(&ControlPacket::RpcResponse { id, result }) - .await?; - } - ControlPacket::RpcResponse { id, result } => session.rpc.resolve(id, result), - ControlPacket::Ping => session.conn.send(&ControlPacket::Pong).await?, - _ => {} - } - Ok(()) -} - -async fn handle_rpc(_session: &Arc, method: &str, args: Vec) -> RpcResult { - match method { - "hello" => RpcResult { - stdout: format!("Hello from server! Args: {:?}", args), - ..Default::default() - }, - _ => RpcResult { - stderr: format!("Unknown method: {}", method), - code: -1, - ..Default::default() - }, - } -} diff --git a/packages/client-v2/src/bin/client.rs b/packages/client-v2/src/bin/client.rs index fdccb94..7850a1c 100644 --- a/packages/client-v2/src/bin/client.rs +++ b/packages/client-v2/src/bin/client.rs @@ -1,6 +1,14 @@ -use xiao::app::client::Client; +#[cfg(target_os = "linux")] use std::sync::Arc; +#[cfg(target_os = "linux")] +use xiao::app::client::Client; +#[cfg(not(target_os = "linux"))] +fn main() { + println!("This client only works on Linux due to ALSA dependencies."); +} + +#[cfg(target_os = "linux")] #[tokio::main] async fn main() -> anyhow::Result<()> { let client = Arc::new(Client::new()); diff --git a/packages/client-v2/src/bin/server.rs b/packages/client-v2/src/bin/server.rs index f6ed176..531f5b9 100644 --- a/packages/client-v2/src/bin/server.rs +++ b/packages/client-v2/src/bin/server.rs @@ -35,7 +35,7 @@ async fn main() -> anyhow::Result<()> { tokio::time::sleep(std::time::Duration::from_secs(12)).await; println!("3. Testing Audio Playback (from temp/test.wav)..."); - server.start_play(addr).await?; + server.start_play(addr, "temp/test.wav").await?; break; } diff --git a/packages/client-v2/src/net/discovery.rs b/packages/client-v2/src/net/discovery.rs index a1b4c8b..a1e1415 100644 --- a/packages/client-v2/src/net/discovery.rs +++ b/packages/client-v2/src/net/discovery.rs @@ -4,22 +4,22 @@ use std::net::{IpAddr, SocketAddr}; use std::time::Duration; use tokio::net::UdpSocket; +const DISCOVERY_PROTOCOL: &str = "XIAO_V2"; + pub const DISCOVERY_PORT: u16 = 53530; -const DISCOVERY_MAGIC: &[u8] = b"XIAO_DISCOVERY_V2"; pub struct Discovery; impl Discovery { - pub async fn broadcast(tcp_port: u16, udp_port: u16) -> Result<()> { + pub async fn broadcast(port: u16) -> Result<()> { let socket = UdpSocket::bind("0.0.0.0:0").await?; socket.set_broadcast(true)?; let target: SocketAddr = format!("255.255.255.255:{}", DISCOVERY_PORT).parse()?; - let mut msg = DISCOVERY_MAGIC.to_vec(); - msg.extend(postcard::to_allocvec(&ControlPacket::ServerHello { - tcp_port, - udp_port, - })?); + let msg = postcard::to_allocvec(&ControlPacket::Discovery { + protocol: DISCOVERY_PROTOCOL.to_string(), + port, + })?; tokio::spawn(async move { loop { @@ -30,19 +30,15 @@ impl Discovery { Ok(()) } - pub async fn listen() -> Result<(IpAddr, u16, u16)> { + pub async fn listen() -> Result<(IpAddr, u16)> { let socket = UdpSocket::bind(format!("0.0.0.0:{}", DISCOVERY_PORT)).await?; let mut buf = [0u8; 1024]; loop { let (len, addr) = socket.recv_from(&mut buf).await?; let data = &buf[..len]; - - if data.starts_with(DISCOVERY_MAGIC) { - let packet_data = &data[DISCOVERY_MAGIC.len()..]; - if let Ok(ControlPacket::ServerHello { tcp_port, udp_port }) = - postcard::from_bytes(packet_data) - { - return Ok((addr.ip(), tcp_port, udp_port)); + if let Ok(ControlPacket::Discovery { protocol, port }) = postcard::from_bytes(data) { + if protocol == DISCOVERY_PROTOCOL { + return Ok((addr.ip(), port)); } } } diff --git a/packages/client-v2/src/net/network.rs b/packages/client-v2/src/net/network.rs index adfbf74..683b504 100644 --- a/packages/client-v2/src/net/network.rs +++ b/packages/client-v2/src/net/network.rs @@ -6,11 +6,6 @@ use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::net::{TcpStream, UdpSocket}; use tokio::sync::Mutex; -pub struct NetConfig { - pub tcp_port: u16, - pub udp_port: u16, -} - /// A unified control connection over TCP pub struct Connection { reader: Mutex, diff --git a/packages/client-v2/src/net/protocol.rs b/packages/client-v2/src/net/protocol.rs index d469e93..8c9a344 100644 --- a/packages/client-v2/src/net/protocol.rs +++ b/packages/client-v2/src/net/protocol.rs @@ -1,46 +1,36 @@ use crate::audio::config::AudioConfig; use serde::{Deserialize, Serialize}; -#[derive(Serialize, Deserialize, Debug, Clone, PartialEq, Eq)] -pub struct DeviceInfo { +#[derive(Serialize, Deserialize, Debug, Clone)] +pub struct ClientInfo { pub model: String, - pub mac: String, - pub version: u32, -} - -impl DeviceInfo { - pub fn current() -> Self { - Self { - model: "Open-XiaoAi-V2".to_string(), - mac: "00:00:00:00:00:00".to_string(), // TODO: Get actual MAC - version: 1, - } - } + pub serial_number: String, } #[derive(Serialize, Deserialize, Debug, Clone)] pub enum ControlPacket { // Discovery - ServerHello { - tcp_port: u16, - udp_port: u16, + Discovery { + protocol: String, + port: u16, }, - // Handshake - ClientIdentify { - info: DeviceInfo, - udp_port: u16, - }, - IdentifyOk, - // Audio Control - StartRecording { - config: AudioConfig, + // Handshake + ServerHello { + auth: String, + version: String, + udp_port: u16, // for audio }, - StopRecording, - StartPlayback { - config: AudioConfig, + ClientHello { + auth: String, + version: String, + udp_port: u16, // for audio + info: ClientInfo, }, - StopPlayback, + + // Heartbeat + Ping, + Pong, // RPC RpcRequest { @@ -53,9 +43,15 @@ pub enum ControlPacket { result: RpcResult, }, - // Heartbeat - Ping, - Pong, + // Audio Control + StartRecording { + config: AudioConfig, + }, + StopRecording, + StartPlayback { + config: AudioConfig, + }, + StopPlayback, } #[derive(Serialize, Deserialize, Debug, Clone, Default)]