From a0a2194ded11b535010dd84970cbb582a5076fc3 Mon Sep 17 00:00:00 2001 From: Claas Date: Sun, 14 Dec 2025 21:37:36 +0100 Subject: [PATCH] Fix race condition --- delivery-service/Cargo.toml | 2 +- delivery-service/src/main.rs | 31 +++++++++++++++++++++---------- 2 files changed, 22 insertions(+), 11 deletions(-) diff --git a/delivery-service/Cargo.toml b/delivery-service/Cargo.toml index 70b8b2e..0e55913 100644 --- a/delivery-service/Cargo.toml +++ b/delivery-service/Cargo.toml @@ -1,7 +1,7 @@ [package] name = "delivery-service" version = "0.1.0" -edition = "2021" +edition = "2024" [dependencies] axum = { workspace = true, features = ["macros", "ws"] } diff --git a/delivery-service/src/main.rs b/delivery-service/src/main.rs index b5fc4fe..c571bc2 100644 --- a/delivery-service/src/main.rs +++ b/delivery-service/src/main.rs @@ -1,19 +1,19 @@ use std::{collections::HashMap, sync::Arc}; use axum::{ + Router, body::Bytes, extract::{ - ws::{self, WebSocket}, Path, State, WebSocketUpgrade, + ws::{self, WebSocket}, }, http::{HeaderName, HeaderValue, StatusCode}, response::IntoResponse, routing::{get, post}, - Router, }; use tokio::{ signal, - sync::{mpsc, Mutex}, + sync::{Mutex, mpsc}, }; use tower_http::{ services::{ServeDir, ServeFile}, @@ -27,7 +27,7 @@ mod extractor; #[derive(Clone, Default)] struct AppState { - channels: Arc, mpsc::Sender>>>, + channels: Arc, Arc>>>>, } /// The single page application setup used in production. During development a vite proxy is used to host the app and @@ -146,6 +146,7 @@ async fn handle_socket(mut socket: WebSocket, State(state): State, cli //TODO decide on buffer size let (sender, mut receiver) = mpsc::channel(8); + let sender = Arc::from(sender); // Register sender for this id // Immediately drop lock after insert to avoid deadlock @@ -153,12 +154,12 @@ async fn handle_socket(mut socket: WebSocket, State(state): State, cli .channels .lock() .await - .insert(client_id.clone(), sender); + .insert(client_id.clone(), sender.clone()); if let Some(previous_sender) = previous_sender { //TODO think about if this is valid tracing::warn!("[{}] Replacing previous subscriber", client_id); - // This should close the SSE stream for the other client that used the same id + // This should close the websocket for the other client that used the same id //TODO test assumption drop(previous_sender); } @@ -169,7 +170,7 @@ async fn handle_socket(mut socket: WebSocket, State(state): State, cli Some(message) = receiver.recv() => { let message = ws::Message::Binary(message.into()); if let Err(error) = socket.send(message).await { - tracing::error!("Error sending message through websocket: {}", error); + tracing::error!("[{}] Error sending message through websocket: {}", client_id, error); //TODO remove channel to avoid memory leak. It is the clients responsibility to reestablish a connection break; } @@ -186,15 +187,25 @@ async fn handle_socket(mut socket: WebSocket, State(state): State, cli // Try to close gracefully but if not ignore error if let Err(error) = socket.close().await { tracing::trace!( - "[{}] Ignoring error closing websocket: {}", + "[{}] Ignoring error from closing disconnected websocket: {}", client_id, error ); } + // It is the clients responsibility to reestablish a new connection + + let mut channels = state.channels.lock().await; // Remove channel to avoid memory leak - // It is the clients responsibility to reestablish a new connection - state.channels.lock().await.remove(&client_id); + if channels + .get(&client_id) + .is_none_or(|current_sender| !Arc::ptr_eq(current_sender, &sender)) + { + tracing::debug!("[{}] already removed or replaced", client_id); + return; + } + + channels.remove(&client_id); tracing::debug!("[{}] removed", client_id); } -- 2.51.2