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

Rust Axum服务处理5000个WebSocket连接遇IO超时问题求助

问题描述

我基于Rust Axum实现了一个WebSocket服务,功能是从Kafka主题读取消息并转发给对应会话的已建立WebSocket连接。在4CPU、8GB内存的服务器上进行5000连接的负载测试时,客户端频繁出现IO超时错误,但相同配置下的Java Spring Boot服务运行正常且延迟极低。核心代码如下:

lazy_static! {
    pub static ref ACTIVE_SESSIONS: DashMap<String, SplitSink<WebSocket, axum::extract::ws::Message>> =
        DashMap::new();
}

#[tokio::main(flavor = "multi_thread", worker_threads = 16)]
let consumer_handles = (0..num_partitions)
        .map(|_| {
            tokio::spawn(run_async_processor(
                broker_url.to_owned(),
                random_group_id.to_owned(),
                topic.to_owned(),
                username.to_owned(),
                password.to_owned(),
                app_config.kafka.consumer.enable_auto_commit.to_owned(),
                offset.to_owned(),
            ))
        })
        .collect::<Vec<_>>();

match TcpListener::bind(&ADDR).await {
        Ok(listener) => {
            let local_addr = listener.local_addr().expect("Failed to get local address");
            tracing::info!("server is running on port {}", local_addr.port());
            axum::serve(listener, app.into_make_service())
                .tcp_nodelay(true)
                .await
                .expect("Server failed to start");
        }
        Err(e) => {
            tracing::error!("Failed to initiate the listener: {}", e);
            std::process::exit(1);
        }
    }
    
// handler function
pub async fn handler(
    ws: WebSocketUpgrade,
    Path(call_id): Path<String>,
    auth_header: Option<TypedHeader<Authorization<Bearer>>>,
    token: Option<Query<TokenQuery>>,
    State(state): State<Arc<AppState>>,
) -> impl IntoResponse {
    let auth_token = match (auth_header, token) {
        (Some(auth_header), _) => auth_header.token().to_string(),
        (None, Some(token)) => token.token.clone(),
        _ => return (StatusCode::UNAUTHORIZED, "auth token is missing").into_response(),
    };

    // Check for valid JWT format
    if !is_valid_jwt(&auth_token) {
        return (StatusCode::UNAUTHORIZED, "invalid JWT format").into_response();
    }

    let account_info =
        match auth_handler::get_account_info(&auth_token, &state.config).await {
            Ok(account_info) => account_info,
            Err(e) => {
                tracing::error!("failed to get the account info {}", e);
                return (StatusCode::UNAUTHORIZED, e.to_string()).into_response();
            }
        };
    let _auth_context = match auth_handler::get_auth_context(
        &account_info,
        &auth_token,
        None,
        &state.auth_properties,
    )
    .await
    {
        Ok(auth_context) => auth_context,
        Err(e) => {
            tracing::error!("failed to authorize the user");
            return (StatusCode::UNAUTHORIZED, e.to_string()).into_response();
        }
    };

    ws.on_upgrade(move |socket| handle_socket(socket, call_id))
}

async fn handle_socket(socket: WebSocket, call_id: String) {
    let (sender, mut receiver) = socket.split();
    {
        ACTIVE_SESSIONS.insert(call_id.clone(), sender);
    }

    while let Some(msg) = receiver.next().await {
        match msg {
            Ok(axum::extract::ws::Message::Text(text)) => {
                match serde_json::from_str::<Value>(&text) {
                    Ok(json) => {
                        match json.get("type").and_then(Value::as_str) {
                            Some("ping") => {
                                tracing::info!("Received ping message");
                                let response = "{\"type\": \"pong\"}";
                                if let Some(mut ws) = ACTIVE_SESSIONS.get_mut(&call_id) {
                                    let msg =
                                        axum::extract::ws::Message::Text(response.to_string());
                                    if let Err(e) = ws.send(msg).await {
                                        tracing::error!(
                                            "error in sending websocket message: {}",
                                            e
                                        );
                                    }
                                    tracing::info!("Sent pong message");
                                }
                            }
                            _ => {
                                tracing::warn!("Unknown message type");
                            }
                        }
                    }
                    Err(e) => {
                        tracing::error!("failed to parse JSON message: {}", e);
                    }
                }
            }
            Ok(axum::extract::ws::Message::Close(frame)) => {
                if let Some(mut ws) = ACTIVE_SESSIONS.get_mut(&call_id) {
                    let msg = axum::extract::ws::Message::Close(frame);
                    if let Err(e) = ws.send(msg).await {
                        tracing::error!("error in sending close message: {}", e);
                    }
                }
                break;
            }
            Ok(_) => {
                continue;
            }
            Err(e) => {
                tracing::error!("WebSocket error for call_id={}: {}", call_id, e);
                break;
            }
        }
    }

    ACTIVE_SESSIONS.remove(&call_id);
    tracing::info!("Removed session for call_id={}", call_id);
}

// kafka consumer function
async fn run_async_processor(
    brokers: String,
    group_id: String,
    topic: String,
    username: String,
    password: String,
    enable_auto_commit: String,
    offset: String,
) {
    let consumer = create_consumer(
        brokers,
        group_id,
        username,
        password,
        enable_auto_commit,
        offset,
    );

    match consumer.subscribe(&[&topic]) {
        Ok(_) => {
            tracing::info!("subscribing to topic={}", topic);
        }
        Err(_e) => {
            panic!("failed to subscribe to the topic={}", topic);
        }
    }

    loop {
        match consumer.recv().await {
            Ok(m) => {
                match m.payload_view::<str>() {
                    None => continue,
                    Some(Ok(payload)) => {
                        let payload = payload.to_owned();

                        tokio::spawn(async move {
                            match serde_json::from_str::<CallTopicMessage>(&payload) {
                                Ok(call_topic_message) => {
                                    tracing::info!(
                                        "received message for call_id={}, type={:?}",
                                        call_topic_message.call_id,
                                        call_topic_message.message_type
                                    );
                                    if let Some(mut ws) =
                                        ACTIVE_SESSIONS.get_mut(&call_topic_message.call_id)
                                    {
                                        tracing::info!(
                                            "active_session for call_id={} found",
                                            &call_topic_message.call_id
                                        );
                                        match serde_json::to_string(
                                            &call_topic_message.agent_assist_message,
                                        ) {
                                            Ok(msg) => {
                                                let msg = axum::extract::ws::Message::Text(msg);
                                                if let Err(e) = ws.send(msg).await {
                                                    tracing::error!(
                                                        "error in sending websocket message: {}",
                                                        e
                                                    );
                                                }
                                            }
                                            Err(e) => {
                                                tracing::error!(
                                                    "error in serializing kafka message {}",
                                                    e
                                                );
                                            }
                                        }
                                    }
                                }
                                Err(e) => {
                                    tracing::error!("error in deserializing kafka message {}", e);
                                }
                            };
                        });
                    }
                    Some(Err(_e)) => {
                    }
                };
            }
            Err(e) => tracing::error!("error in consuming from Kafka: {}", e),
        }
    }
}
优化建议

1. 修复DashMap的异步锁持有问题

当前用DashMap存储SplitSink,调用get_mut后会持有全局锁直到send异步操作完成,高并发下会导致其他线程阻塞,严重拖慢性能。

  • 优化方案:将SplitSink包装为Arc<Mutex<SplitSink>>,把全局锁拆分为每个连接的细粒度锁,减少全局锁持有时间。
  • 代码调整:
    lazy_static! {
        pub static ref ACTIVE_SESSIONS: DashMap<String, Arc<Mutex<SplitSink<WebSocket, axum::extract::ws::Message>>>> =
            DashMap::new();
    }
    
    // 插入连接时
    ACTIVE_SESSIONS.insert(call_id.clone(), Arc::new(Mutex::new(sender)));
    
    // 发送消息时
    if let Some(session) = ACTIVE_SESSIONS.get(&call_id) {
        let mut ws = session.value().lock().await;
        let send_result = ws.send(msg).await;
        // 处理发送错误
    }
    

2. 限制Kafka消息处理的并发数

当前每收到一条Kafka消息就创建一个新Tokio任务,高消息量下会导致任务爆炸,引发调度延迟。

  • 优化方案:用tokio::sync::Semaphore限制并发处理的任务数,比如根据CPU核心数设置为8或16。
  • 代码示例:
    // 在consumer初始化时创建Semaphore
    let semaphore = Arc::new(Semaphore::new(8));
    
    // 处理消息时
    let permit = semaphore.acquire().await.unwrap();
    tokio::spawn(async move {
        // 消息处理逻辑
        drop(permit); // 释放许可
    });
    

3. 调整Tokio运行时配置

当前设置worker_threads = 16,但4CPU机器上过多线程会增大上下文切换开销。

  • 优化方案:将工作线程数改为CPU核心数的2倍(即8),或者直接移除手动配置,让Tokio自动适配硬件。
  • 同时检查是否存在阻塞操作(如JWT验证、同步IO),如果有,用tokio::task::spawn_blocking将其移到阻塞线程池,避免占用工作线程。

4. 优化WebSocket连接清理逻辑

当前仅在连接循环退出时移除会话,若send操作失败(如连接断开),无效连接会留在ACTIVE_SESSIONS中,浪费资源。

  • 优化方案:在send失败时直接移除对应会话:
    if let Err(e) = ws.send(msg).await {
        tracing::error!("error in sending websocket message: {}", e);
        ACTIVE_SESSIONS.remove(&call_id);
    }
    

5. 调优TCP参数

  • 增大TCP监听队列的backlog,避免高并发下连接溢出:
    let listener = TcpListener::bind(&ADDR).await?;
    listener.set_backlog(4096)?; // 根据实际情况调整
    
  • 开启TCP keepalive,维持空闲连接存活:
    axum::serve(listener, app.into_make_service())
        .tcp_nodelay(true)
        .tcp_keepalive(Some(TcpKeepalive::default()))
        .await?;
    

6. 减少日志输出开销

当前大量info级日志在高并发下会成为IO瓶颈,建议将非关键日志改为debug级别,生产环境仅保留error和warn级别。

内容的提问来源于stack exchange,提问作者Sumit Kumar

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 09:00:53