chore: polish client/server code
This commit is contained in:
@@ -11,7 +11,7 @@ use tokio_util::sync::CancellationToken;
|
||||
|
||||
pub struct ClientAudioManager {
|
||||
audio_socket: Arc<AudioSocket>,
|
||||
server_audio_addr: SocketAddr,
|
||||
audio_addr: SocketAddr,
|
||||
session_cancel: CancellationToken,
|
||||
record_cancel: RwLock<Option<CancellationToken>>,
|
||||
play_cancel: RwLock<Option<CancellationToken>>,
|
||||
@@ -20,12 +20,12 @@ pub struct ClientAudioManager {
|
||||
impl ClientAudioManager {
|
||||
pub fn new(
|
||||
audio_socket: Arc<AudioSocket>,
|
||||
server_audio_addr: SocketAddr,
|
||||
audio_addr: SocketAddr,
|
||||
session_cancel: CancellationToken,
|
||||
) -> Self {
|
||||
Self {
|
||||
audio_socket,
|
||||
server_audio_addr,
|
||||
audio_addr,
|
||||
session_cancel,
|
||||
record_cancel: RwLock::new(None),
|
||||
play_cancel: RwLock::new(None),
|
||||
@@ -52,7 +52,7 @@ impl ClientAudioManager {
|
||||
*self.record_cancel.write().await = Some(token.clone());
|
||||
|
||||
let audio_socket = self.audio_socket.clone();
|
||||
let server_audio_addr = self.server_audio_addr;
|
||||
let audio_addr = self.audio_addr;
|
||||
|
||||
tokio::spawn(async move {
|
||||
let (pcm_tx, mut pcm_rx) = mpsc::channel::<Vec<i16>>(32);
|
||||
@@ -89,7 +89,7 @@ impl ClientAudioManager {
|
||||
Some(pcm) = pcm_rx.recv() => {
|
||||
let mut out = vec![0u8; 4096];
|
||||
if let Ok(len) = codec.encode(&pcm, &mut out) {
|
||||
let _ = audio_socket.send(&AudioPacket { data: out[..len].to_vec() }, server_audio_addr).await;
|
||||
let _ = audio_socket.send(&AudioPacket { data: out[..len].to_vec() }, audio_addr).await;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -10,14 +10,13 @@ use anyhow::{Result, anyhow};
|
||||
use audio_manager::ClientAudioManager;
|
||||
use std::net::SocketAddr;
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::{RwLock, mpsc};
|
||||
use tokio::sync::{RwLock};
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
/// 内部连接上下文,包含了音频流所需的全部信息
|
||||
struct ActiveSession {
|
||||
conn: Arc<Connection>,
|
||||
audio_manager: ClientAudioManager,
|
||||
session_cancel: CancellationToken, // 控制整个 Session 的生命周期
|
||||
session_cancel: CancellationToken,
|
||||
}
|
||||
|
||||
pub struct Client {
|
||||
@@ -35,11 +34,12 @@ impl Client {
|
||||
|
||||
pub async fn run(self: Arc<Self>) -> Result<()> {
|
||||
loop {
|
||||
println!("Searching for server...");
|
||||
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;
|
||||
tokio::time::sleep(std::time::Duration::from_secs(1)).await;
|
||||
continue;
|
||||
}
|
||||
};
|
||||
@@ -47,20 +47,17 @@ impl Client {
|
||||
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);
|
||||
match tokio::net::TcpStream::connect(addr).await {
|
||||
Err(e) => eprintln!("Failed to connect to {}: {}", addr, e),
|
||||
Ok(stream) => {
|
||||
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;
|
||||
}
|
||||
}
|
||||
@@ -83,14 +80,15 @@ impl Client {
|
||||
// --- 握手 (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());
|
||||
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(|_| "open-xiaoai".to_string());
|
||||
std::env::var("XIAO_CLIENT_AUTH").unwrap_or_else(|_| "xiao-client".to_string());
|
||||
|
||||
conn.send(&ControlPacket::ClientHello {
|
||||
auth: server_auth,
|
||||
version: version.clone(),
|
||||
udp_port: audio_socket.port(),
|
||||
// todo client info
|
||||
info: ClientInfo {
|
||||
model: "Open-XiaoAi-V2".to_string(),
|
||||
serial_number: "00:00:00:00:00:00".to_string(),
|
||||
@@ -115,10 +113,10 @@ impl Client {
|
||||
_ => return Err(anyhow!("Handshake failed")),
|
||||
};
|
||||
|
||||
let server_audio_addr = SocketAddr::new(addr.ip(), server_udp_port);
|
||||
let audio_addr = SocketAddr::new(addr.ip(), server_udp_port);
|
||||
println!(
|
||||
"Handshake successful with {}, audio at {}",
|
||||
addr, server_audio_addr
|
||||
addr, audio_addr
|
||||
);
|
||||
|
||||
// --- 初始化 Session ---
|
||||
@@ -127,17 +125,14 @@ impl Client {
|
||||
conn: conn.clone(),
|
||||
audio_manager: ClientAudioManager::new(
|
||||
audio_socket,
|
||||
server_audio_addr,
|
||||
audio_addr,
|
||||
session_cancel.clone(),
|
||||
),
|
||||
session_cancel,
|
||||
});
|
||||
*self.session.write().await = Some(session.clone());
|
||||
|
||||
// 使用 mpsc 队列来缓冲指令,确保顺序执行且不阻塞接收循环
|
||||
let (cmd_tx, mut cmd_rx) = mpsc::channel::<ControlPacket>(64);
|
||||
|
||||
// 任务 1: 心跳
|
||||
// 心跳
|
||||
let hb_session = session.clone();
|
||||
tokio::spawn(async move {
|
||||
let mut interval = tokio::time::interval(std::time::Duration::from_secs(10));
|
||||
@@ -151,25 +146,16 @@ impl Client {
|
||||
}
|
||||
});
|
||||
|
||||
// 任务 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; }
|
||||
if let Err(e) = self.process_packet(packet, &session).await {
|
||||
eprintln!("Process packet error: {}", e);
|
||||
}
|
||||
}
|
||||
Ok(Err(e)) => {
|
||||
return Err(anyhow!("Connection receive error: {}", e));
|
||||
@@ -183,6 +169,7 @@ impl Client {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// todo 考虑有些操作比较耗时,需要非阻塞处理
|
||||
async fn process_packet(
|
||||
&self,
|
||||
packet: ControlPacket,
|
||||
|
||||
@@ -29,10 +29,17 @@ pub struct ServerAudioManager {
|
||||
}
|
||||
|
||||
impl ServerAudioManager {
|
||||
pub fn new(session_cancel: CancellationToken, tracker: TaskTracker) -> (Self, mpsc::Receiver<AudioPacket>, mpsc::Receiver<RecorderCommand>) {
|
||||
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),
|
||||
@@ -41,7 +48,7 @@ impl ServerAudioManager {
|
||||
recorder_tx,
|
||||
tracker,
|
||||
};
|
||||
|
||||
|
||||
(manager, audio_rx, recorder_rx)
|
||||
}
|
||||
|
||||
@@ -60,10 +67,7 @@ impl ServerAudioManager {
|
||||
}
|
||||
|
||||
self.recorder_tx
|
||||
.send(RecorderCommand::Start {
|
||||
config,
|
||||
filename,
|
||||
})
|
||||
.send(RecorderCommand::Start { config, filename })
|
||||
.await?;
|
||||
Ok(())
|
||||
}
|
||||
@@ -94,7 +98,7 @@ impl ServerAudioManager {
|
||||
};
|
||||
|
||||
let session_cancel = self.session_cancel.clone();
|
||||
|
||||
|
||||
self.tracker.spawn(async move {
|
||||
let mut reader = reader;
|
||||
let mut codec = match OpusCodec::new(&config) {
|
||||
@@ -146,7 +150,11 @@ impl ServerAudioManager {
|
||||
}
|
||||
}
|
||||
|
||||
pub fn spawn_audio_processor(&self, mut audio_rx: mpsc::Receiver<AudioPacket>, mut recorder_rx: mpsc::Receiver<RecorderCommand>) {
|
||||
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;
|
||||
@@ -198,4 +206,3 @@ impl ServerAudioManager {
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -6,8 +6,8 @@ use crate::net::discovery::Discovery;
|
||||
use crate::net::network::{AudioSocket, Connection};
|
||||
use crate::net::protocol::{ClientInfo, ControlPacket, RpcResult};
|
||||
use crate::net::rpc::RpcManager;
|
||||
use audio_manager::ServerAudioManager;
|
||||
use anyhow::{Context, Result, anyhow};
|
||||
use audio_manager::ServerAudioManager;
|
||||
use dashmap::DashMap;
|
||||
use std::net::SocketAddr;
|
||||
use std::sync::Arc;
|
||||
@@ -72,7 +72,7 @@ impl Server {
|
||||
let (stream, addr) = listener.accept().await?;
|
||||
let server = self.clone();
|
||||
tokio::spawn(async move {
|
||||
if let Err(e) = server.clone().handle_connection(stream, addr).await {
|
||||
if let Err(e) = server.clone().handle_session(stream, addr).await {
|
||||
eprintln!("Session {} error: {}", addr, e);
|
||||
}
|
||||
server.remove_session(&addr).await;
|
||||
@@ -89,7 +89,7 @@ impl Server {
|
||||
}
|
||||
}
|
||||
|
||||
async fn handle_connection(
|
||||
async fn handle_session(
|
||||
self: Arc<Self>,
|
||||
stream: tokio::net::TcpStream,
|
||||
addr: SocketAddr,
|
||||
@@ -100,9 +100,9 @@ impl Server {
|
||||
// --- 握手 (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());
|
||||
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(|_| "open-xiaoai".to_string());
|
||||
std::env::var("XIAO_CLIENT_AUTH").unwrap_or_else(|_| "xiao-client".to_string());
|
||||
|
||||
let (info, client_audio_port) = match conn.recv().await? {
|
||||
ControlPacket::ClientHello {
|
||||
@@ -119,9 +119,7 @@ impl Server {
|
||||
}
|
||||
(info, udp_port)
|
||||
}
|
||||
p => {
|
||||
return Err(anyhow!("Handshake failed: unexpected packet {:?}", p));
|
||||
}
|
||||
_ => return Err(anyhow!("Handshake failed")),
|
||||
};
|
||||
|
||||
conn.send(&ControlPacket::ServerHello {
|
||||
@@ -137,6 +135,7 @@ impl Server {
|
||||
info.model, info.serial_number, audio_addr
|
||||
);
|
||||
|
||||
// --- 初始化 Session ---
|
||||
let tracker = TaskTracker::new();
|
||||
let session_cancel = CancellationToken::new();
|
||||
let (audio_manager, audio_rx, recorder_rx) =
|
||||
@@ -153,20 +152,17 @@ impl Server {
|
||||
tracker: tracker.clone(),
|
||||
});
|
||||
|
||||
session.audio_manager.spawn_audio_processor(audio_rx, recorder_rx);
|
||||
session
|
||||
.audio_manager
|
||||
.spawn_audio_processor(audio_rx, recorder_rx);
|
||||
|
||||
self.sessions.insert(addr, session.clone());
|
||||
self.udp_to_tcp.insert(audio_addr, addr);
|
||||
|
||||
// --- 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)) => {
|
||||
@@ -190,6 +186,11 @@ 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<()> {
|
||||
match packet {
|
||||
ControlPacket::Ping => {
|
||||
@@ -211,8 +212,23 @@ impl Server {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn get_clients(&self) -> Vec<SocketAddr> {
|
||||
self.sessions.iter().map(|r| *r.key()).collect()
|
||||
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()
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn call(
|
||||
@@ -238,25 +254,6 @@ impl Server {
|
||||
Ok(rx.await?)
|
||||
}
|
||||
|
||||
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()
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn start_record(&self, addr: SocketAddr, config: AudioConfig) -> Result<()> {
|
||||
let session = self
|
||||
.sessions
|
||||
@@ -324,12 +321,10 @@ impl Server {
|
||||
})
|
||||
.await?;
|
||||
|
||||
session.audio_manager.start_playback(
|
||||
config,
|
||||
reader,
|
||||
self.audio.clone(),
|
||||
session.audio_addr,
|
||||
).await?;
|
||||
session
|
||||
.audio_manager
|
||||
.start_playback(config, reader, self.audio.clone(), session.audio_addr)
|
||||
.await?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user