You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何修改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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.24 15:55:02