如何修改tokio-tungstenite的accept_connection函数接入广播接收器?
问题描述
我正在开发一个WebSocket服务器,需要实现以下功能:
- 将收到的WebSocket消息转发至队列(已完成);
- 通过WebSocket发送来自另一个队列的消息(待解决)。
后台任务会从队列读取消息进行后续处理,为便于提问已简化应用。我先尝试使用tokio::sync::mpsc,但发现接收器无法在多个任务间共享;转而使用tokio::sync::broadcast通道,却无法成功编译代码。我计划使用tokio::select!宏处理收发消息,请问该如何修改accept_connection函数以访问接收器?
可运行的初始代码
use std::{env, io::Error}; use futures_util::StreamExt; use tokio::net::{TcpListener, TcpStream}; use tokio::sync::broadcast; use tokio::sync::mpsc; #[tokio::main] async fn main() -> Result<(), Error> { let addr = env::args().nth(1).unwrap_or_else(|| "127.0.0.1:8080".to_string()); // 创建事件循环和TCP监听器 let try_socket = TcpListener::bind(&addr).await; let listener = try_socket.expect("绑定失败"); // 初始化通道 let (tx_from_ws, mut rx_from_ws) = mpsc::channel::<String>(32); let (tx_to_ws, mut rx_to_ws) = broadcast::channel::<String>(32); // 打印所有从WebSocket收到的消息 tokio::spawn(async move { while let Some(msg) = rx_from_ws.recv().await { println!("From websocket: {}", msg); } }); // 定期向所有WebSocket连接发送消息 tokio::spawn(async move { let mut interval = tokio::time::interval(tokio::time::Duration::from_secs(5)); loop { interval.tick().await; tx_to_ws.send("Periodic message".to_string()).expect("发送消息到通道失败"); } }); while let Ok((stream, _)) = listener.accept().await { tokio::spawn(accept_connection(stream, tx_from_ws.clone())); } Ok(()) } async fn accept_connection(stream: TcpStream, tx: mpsc::Sender<String>) { let addr = stream.peer_addr().expect("连接的流必须有对等地址"); let ws_stream = tokio_tungstenite::accept_async(stream) .await .expect("WebSocket握手出错"); let (write, mut read) = ws_stream.split(); loop { tokio::select! { Some(msg) = read.next() => { let msg = msg.expect("读取WebSocket消息出错"); let msg = msg.to_text().expect("转换消息为文本出错"); tx.send(msg.to_string()).await.expect("发送消息到通道失败"); } // 待处理来自通道的消息 } } }
Cargo.toml配置
[package] name = "rust-websocket-test" version = "0.1.0" edition = "2021" [dependencies] tokio = { version = "1.27.0", features = ["full"] } tokio-tungstenite = "*" futures-util = "0.3.17" futures-channel = "0.3.17"
尝试实现接收器后的代码
use std::{env, io::Error}; use futures_util::StreamExt; use tokio::net::{TcpListener, TcpStream}; use tokio::sync::broadcast; use tokio::sync::mpsc; #[tokio::main] async fn main() -> Result<(), Error> { let addr = env::args().nth(1).unwrap_or_else(|| "127.0.0.1:8080".to_string()); // 创建事件循环和TCP监听器 let try_socket = TcpListener::bind(&addr).await; let listener = try_socket.expect("绑定失败"); // 初始化通道 let (tx_from_ws, mut rx_from_ws) = mpsc::channel::<String>(32); let (tx_to_ws, mut rx_to_ws) = broadcast::channel::<String>(32); // 尝试用mutex包装rx_to_ws // 打印所有从WebSocket收到的消息 tokio::spawn(async move { while let Some(msg) = rx_from_ws.recv().await { println!("From websocket: {}", msg); } }); // 定期向所有WebSocket连接发送消息 tokio::spawn(async move { let mut interval = tokio::time::interval(tokio::time::Duration::from_secs(5)); loop { interval.tick().await; tx_to_ws.send("Periodic message".to_string()).expect("发送消息到通道失败"); } }); while let Ok((stream, _)) = listener.accept().await { tokio::spawn(accept_connection(stream, tx_from_ws.clone(), rx_to_ws.clone())); } Ok(()) } async fn accept_connection(stream: TcpStream, tx: mpsc::Sender<String>, rx: broadcast::Receiver<String>) { let addr = stream.peer_addr().expect("连接的流必须有对等地址"); let ws_stream = tokio_tungstenite::accept_async(stream) .await .expect("WebSocket握手出错"); let (write, mut read) = ws_stream.split(); loop { tokio::select! { Some(msg) = read.next() => { let msg = msg.expect("读取WebSocket消息出错"); let msg = msg.to_text().expect("转换消息为文本出错"); tx.send(msg.to_string()).await.expect("发送消息到通道失败"); } // 处理来自通道的消息 Some(msg) = rx.recv() => { let msg = tokio_tungstenite::tungstenite::Message::text(msg); tokio::spawn(async move { if let Err(e) = write.send(msg).await { eprintln!("发送消息到WebSocket出错: {}", e); } }); } } } }
解决方案
你的代码存在两个核心问题,修改后即可正常编译运行:
1. 修复Broadcast接收器的调用逻辑
broadcast::Receiver的recv()方法返回Result<String, broadcast::RecvError>而非Option,不能用Some(msg) = rx.recv()的写法,必须处理通道关闭、消息滞后等错误场景。
2. 解决WebSocket写入端的所有权问题
你将write转移到新的tokio::spawn任务后,后续循环无法再使用该写入端。需要用Arc<Mutex<WriteHalf>>包装写入端,实现多任务安全共享。
修改后的完整代码如下:
use std::{env, io::Error, sync::Arc}; use futures_util::{StreamExt, SinkExt}; use tokio::net::{TcpListener, TcpStream}; use tokio::sync::{broadcast, mpsc, Mutex}; #[tokio::main] async fn main() -> Result<(), Error> { let addr = env::args().nth(1).unwrap_or_else(|| "127.0.0.1:8080".to_string()); let try_socket = TcpListener::bind(&addr).await; let listener = try_socket.expect("绑定失败"); let (tx_from_ws, mut rx_from_ws) = mpsc::channel::<String>(32); // 移除rx_to_ws的mut,因为我们要克隆给每个连接 let (tx_to_ws, rx_to_ws) = broadcast::channel::<String>(32); tokio::spawn(async move { while let Some(msg) = rx_from_ws.recv().await { println!("From websocket: {}", msg); } }); tokio::spawn(async move { let mut interval = tokio::time::interval(tokio::time::Duration::from_secs(5)); loop { interval.tick().await; tx_to_ws.send("Periodic message".to_string()).expect("发送消息到通道失败"); } }); while let Ok((stream, _)) = listener.accept().await { tokio::spawn(accept_connection(stream, tx_from_ws.clone(), rx_to_ws.clone())); } Ok(()) } async fn accept_connection(stream: TcpStream, tx: mpsc::Sender<String>, mut rx: broadcast::Receiver<String>) { let addr = stream.peer_addr().expect("连接的流必须有对等地址"); let ws_stream = tokio_tungstenite::accept_async(stream) .await .expect("WebSocket握手出错"); // 用Arc<Mutex>包装写入端,实现多任务共享 let (write, mut read) = ws_stream.split(); let write = Arc::new(Mutex::new(write)); loop { tokio::select! { msg_result = read.next() => { match msg_result { Some(Ok(msg)) => { let msg_text = msg.to_text().expect("转换消息为文本出错"); tx.send(msg_text.to_string()).await.expect("发送到mpsc通道失败"); } Some(Err(e)) => { eprintln!("{}的WebSocket读取错误: {}", addr, e); break; // 连接出错,退出循环 } None => { println!("{}的WebSocket连接已关闭", addr); break; // 连接关闭,退出循环 } } } recv_result = rx.recv() => { match recv_result { Ok(msg) => { let ws_msg = tokio_tungstenite::tungstenite::Message::text(msg); let write_clone = write.clone(); tokio::spawn(async move { if let Err(e) = write_clone.lock().await.send(ws_msg).await { eprintln!("向{}发送消息失败: {}", addr, e); } }); } Err(broadcast::RecvError::Closed) => { eprintln!("{}的Broadcast通道已关闭", addr); break; } Err(broadcast::RecvError::Lagged(count)) => { eprintln!("{}丢失了{}条消息", addr, count); } } } } } }
修改说明
- Broadcast错误处理:正确解析
recv()返回的Result,处理通道关闭、消息滞后等情况。 - 写入端共享:通过
Arc<Mutex>包装WebSocket写入端,避免所有权转移导致的编译错误。 - 完善错误处理:为WebSocket读写添加错误分支,连接异常时主动退出循环,避免无限挂起。
内容的提问来源于stack exchange,提问作者Trantidon
相关产品推荐
相关产品推荐

