// file: crates/common/game-realtime-websocket-lib/src/websocket.rs // version: 2 use futures_util::SinkExt; // rust-rules: trait-import use futures_util::StreamExt; // rust-rules: trait-import const TRACING_TARGET: &str = "games::realtime::websocket"; type ClientStream = tokio_tungstenite::WebSocketStream>; type ServerStream = tokio_tungstenite::WebSocketStream; type ClientSink = futures_util::stream::SplitSink; type ServerSink = futures_util::stream::SplitSink; type ClientReceiver = futures_util::stream::SplitStream; type ServerReceiver = futures_util::stream::SplitStream; enum WebSocketStreamKind { Client(ClientStream), Server(ServerStream), } enum WebSocketSinkKind { Client(ClientSink), Server(ServerSink), } enum WebSocketReceiverKind { Client(ClientReceiver), Server(ServerReceiver), } /// Established WebSocket connection implementing the transport-neutral realtime contract. pub struct WebSocketConnection { inner: WebSocketStreamKind, config: crate::WebSocketConfig, } impl WebSocketConnection { fn from_client(stream: ClientStream, config: crate::WebSocketConfig) -> Self { return Self { inner: WebSocketStreamKind::Client(stream), config }; } fn from_server(stream: ServerStream, config: crate::WebSocketConfig) -> Self { return Self { inner: WebSocketStreamKind::Server(stream), config }; } } impl game_realtime_transport_lib::RealtimeConnection for WebSocketConnection { type Sender = crate::WebSocketSender; type Receiver = crate::WebSocketReceiver; fn split(self) -> (Self::Sender, Self::Receiver) { let max_frame_size = self.config.max_frame_size(); let max_message_size = self.config.max_message_size(); let send_timeout = self.config.send_timeout(); let close_timeout = self.config.close_timeout(); return match self.inner { WebSocketStreamKind::Client(stream) => { let (sender, receiver) = stream.split(); ( crate::WebSocketSender { inner: WebSocketSinkKind::Client(sender), max_frame_size, max_message_size, send_timeout, close_timeout, send_timed_out: false, }, crate::WebSocketReceiver { inner: WebSocketReceiverKind::Client(receiver) }, ) }, WebSocketStreamKind::Server(stream) => { let (sender, receiver) = stream.split(); ( crate::WebSocketSender { inner: WebSocketSinkKind::Server(sender), max_frame_size, max_message_size, send_timeout, close_timeout, send_timed_out: false, }, crate::WebSocketReceiver { inner: WebSocketReceiverKind::Server(receiver) }, ) }, }; } } /// Send half of an established WebSocket transport connection. pub struct WebSocketSender { inner: WebSocketSinkKind, max_frame_size: usize, max_message_size: usize, send_timeout: std::time::Duration, close_timeout: std::time::Duration, send_timed_out: bool, } impl game_realtime_transport_lib::RealtimeSender for WebSocketSender { type SendFuture<'a> = futures_util::future::LocalBoxFuture<'a, Result<(), game_realtime_transport_lib::TransportError>> where Self: 'a; type CloseFuture<'a> = futures_util::future::LocalBoxFuture<'a, Result<(), game_realtime_transport_lib::TransportError>> where Self: 'a; fn send(&mut self, message: game_realtime_transport_lib::TransportMessage) -> Self::SendFuture<'_> { return Box::pin(async move { if self.send_timed_out { return Err(game_realtime_transport_lib::TransportError::new( game_realtime_transport_lib::TransportErrorKind::Aborted, "sender is unavailable after a previous send timeout", )); } let payload_len = message.len(); if payload_len > self.max_message_size || payload_len > self.max_frame_size { let error = game_realtime_transport_lib::TransportError::new( game_realtime_transport_lib::TransportErrorKind::MessageTooLarge, format!("binary payload size {payload_len} exceeds configured message/frame maxima {}/{}", self.max_message_size, self.max_frame_size), ); tracing::warn!( target: TRACING_TARGET, payload_len = payload_len, max_message_size = self.max_message_size, max_frame_size = self.max_frame_size, "outbound WebSocket payload rejected" ); return Err(error); } let websocket_message = tokio_tungstenite::tungstenite::Message::Binary(message.into_bytes().into()); let send_timeout = self.send_timeout; let send = async { return match &mut self.inner { WebSocketSinkKind::Client(sender) => sender.send(websocket_message).await, WebSocketSinkKind::Server(sender) => sender.send(websocket_message).await, }; }; let result = tokio::time::timeout(send_timeout, send).await; return match result { Ok(Ok(())) => { tracing::trace!(target: TRACING_TARGET, payload_len = payload_len, "binary WebSocket payload sent"); Ok(()) }, Ok(Err(error)) => { let mapped = map_stream_error(error); tracing::warn!(target: TRACING_TARGET, kind = %mapped.kind(), detail = mapped.detail(), "WebSocket send failed"); Err(mapped) }, Err(_) => { self.send_timed_out = true; let error = timeout_error("WebSocket send", send_timeout); tracing::warn!(target: TRACING_TARGET, timeout_ms = duration_millis(send_timeout), "WebSocket send timed out"); Err(error) }, }; }); } fn close(&mut self) -> Self::CloseFuture<'_> { return Box::pin(async move { let close_timeout = self.close_timeout; let close = async { return match &mut self.inner { WebSocketSinkKind::Client(sender) => sender.close().await, WebSocketSinkKind::Server(sender) => sender.close().await, }; }; let result = tokio::time::timeout(close_timeout, close).await; return match result { Ok(Ok(())) => { tracing::debug!(target: TRACING_TARGET, "local WebSocket close initiated"); Ok(()) }, Ok(Err(error)) => { let mapped = map_stream_error(error); tracing::warn!(target: TRACING_TARGET, kind = %mapped.kind(), detail = mapped.detail(), "WebSocket close failed"); Err(mapped) }, Err(_) => { let error = timeout_error("WebSocket close", close_timeout); tracing::warn!(target: TRACING_TARGET, timeout_ms = duration_millis(close_timeout), "WebSocket close timed out"); Err(error) }, }; }); } } /// Receive half of an established WebSocket transport connection. pub struct WebSocketReceiver { inner: WebSocketReceiverKind, } impl game_realtime_transport_lib::RealtimeReceiver for WebSocketReceiver { type ReceiveFuture<'a> = futures_util::future::LocalBoxFuture<'a, Result> where Self: 'a; fn receive(&mut self) -> Self::ReceiveFuture<'_> { return Box::pin(async move { loop { let next_message = match &mut self.inner { WebSocketReceiverKind::Client(receiver) => receiver.next().await, WebSocketReceiverKind::Server(receiver) => receiver.next().await, }; match next_message { Some(Ok(tokio_tungstenite::tungstenite::Message::Binary(bytes))) => { tracing::trace!(target: TRACING_TARGET, payload_len = bytes.len(), "binary WebSocket payload received"); return Ok(game_realtime_transport_lib::TransportReceive::Message(game_realtime_transport_lib::TransportMessage::new(bytes.to_vec()))); }, Some(Ok(tokio_tungstenite::tungstenite::Message::Close(_))) => { tracing::debug!(target: TRACING_TARGET, "remote WebSocket close observed"); return Ok(game_realtime_transport_lib::TransportReceive::Closed); }, Some(Ok(tokio_tungstenite::tungstenite::Message::Ping(_))) | Some(Ok(tokio_tungstenite::tungstenite::Message::Pong(_))) => {}, Some(Ok(tokio_tungstenite::tungstenite::Message::Text(_))) => { let error = game_realtime_transport_lib::TransportError::new( game_realtime_transport_lib::TransportErrorKind::Protocol, "text WebSocket messages are not part of the binary transport contract", ); tracing::warn!(target: TRACING_TARGET, kind = %error.kind(), "unsupported WebSocket text message received"); return Err(error); }, Some(Ok(tokio_tungstenite::tungstenite::Message::Frame(_))) => { let error = game_realtime_transport_lib::TransportError::new( game_realtime_transport_lib::TransportErrorKind::Protocol, "unexpected raw WebSocket frame surfaced by the backend", ); tracing::warn!(target: TRACING_TARGET, kind = %error.kind(), "unexpected raw WebSocket frame received"); return Err(error); }, Some(Err(error)) => { let mapped = map_stream_error(error); tracing::warn!(target: TRACING_TARGET, kind = %mapped.kind(), detail = mapped.detail(), "WebSocket receive failed"); return Err(mapped); }, None => { tracing::debug!(target: TRACING_TARGET, "WebSocket stream ended"); return Ok(game_realtime_transport_lib::TransportReceive::Closed); }, } } }); } } /// Bound TCP listener that upgrades accepted peers to WebSocket connections. pub struct WebSocketListener { listener: tokio::net::TcpListener, local_addr: std::net::SocketAddr, config: crate::WebSocketConfig, } impl WebSocketListener { /// Binds a WebSocket listener with the baseline configuration. pub async fn bind(address: std::net::SocketAddr) -> Result { return Self::bind_with_config(address, crate::WebSocketConfig::default()).await; } /// Binds a WebSocket listener with explicit product-facing limits and deadlines. pub async fn bind_with_config(address: std::net::SocketAddr, config: crate::WebSocketConfig) -> Result { if let Err(error) = config.validate() { return Err(error); } let listener = match tokio::net::TcpListener::bind(address).await { Ok(value) => value, Err(error) => { let mapped = game_realtime_transport_lib::TransportError::new(game_realtime_transport_lib::TransportErrorKind::Bind, error.to_string()); tracing::warn!(target: TRACING_TARGET, address = %address, detail = mapped.detail(), "WebSocket listener bind failed"); return Err(mapped); }, }; let local_addr = match listener.local_addr() { Ok(value) => value, Err(error) => { let mapped = game_realtime_transport_lib::TransportError::new(game_realtime_transport_lib::TransportErrorKind::Bind, error.to_string()); tracing::warn!(target: TRACING_TARGET, detail = mapped.detail(), "bound WebSocket listener address lookup failed"); return Err(mapped); }, }; tracing::info!(target: TRACING_TARGET, address = %local_addr, "WebSocket listener bound"); return Ok(Self { listener, local_addr, config }); } /// Returns the concrete local socket address, including an ephemeral port selected by the OS. #[must_use] pub fn local_addr(&self) -> std::net::SocketAddr { return self.local_addr; } /// Accepts one TCP peer and completes a bounded server-side WebSocket handshake. pub async fn accept(&self) -> Result { let (stream, peer_addr) = match self.listener.accept().await { Ok(value) => value, Err(error) => { let mapped = game_realtime_transport_lib::TransportError::new(game_realtime_transport_lib::TransportErrorKind::Accept, error.to_string()); tracing::warn!(target: TRACING_TARGET, detail = mapped.detail(), "WebSocket TCP accept failed"); return Err(mapped); }, }; let handshake = tokio_tungstenite::accept_async_with_config(stream, Some(tungstenite_config(&self.config))); let result = tokio::time::timeout(self.config.connect_timeout(), handshake).await; let websocket = match result { Ok(Ok(value)) => value, Ok(Err(error)) => { let mapped = game_realtime_transport_lib::TransportError::new(game_realtime_transport_lib::TransportErrorKind::Accept, error.to_string()); tracing::warn!(target: TRACING_TARGET, peer = %peer_addr, detail = mapped.detail(), "WebSocket server handshake failed"); return Err(mapped); }, Err(_) => { let error = timeout_error("WebSocket server handshake", self.config.connect_timeout()); tracing::warn!(target: TRACING_TARGET, peer = %peer_addr, timeout_ms = duration_millis(self.config.connect_timeout()), "WebSocket server handshake timed out"); return Err(error); }, }; tracing::info!(target: TRACING_TARGET, peer = %peer_addr, "WebSocket peer accepted"); return Ok(crate::WebSocketConnection::from_server(websocket, self.config)); } } /// Connects a client to one plain `ws://` endpoint with the baseline configuration. pub async fn connect(endpoint: &str) -> Result { return crate::connect_with_config(endpoint, crate::WebSocketConfig::default()).await; } /// Connects a client to one plain `ws://` endpoint with explicit limits and deadlines. pub async fn connect_with_config( endpoint: &str, config: crate::WebSocketConfig, ) -> Result { if !endpoint.starts_with("ws://") { return Err(game_realtime_transport_lib::TransportError::new( game_realtime_transport_lib::TransportErrorKind::InvalidConfiguration, "the baseline WebSocket backend accepts only ws:// endpoints", )); } if let Err(error) = config.validate() { return Err(error); } tracing::debug!(target: TRACING_TARGET, endpoint = endpoint, "connecting WebSocket client"); let handshake = tokio_tungstenite::connect_async_with_config(endpoint, Some(tungstenite_config(&config)), false); let result = tokio::time::timeout(config.connect_timeout(), handshake).await; let (stream, _) = match result { Ok(Ok(value)) => value, Ok(Err(error)) => { let mapped = map_connect_error(error); tracing::warn!( target: TRACING_TARGET, endpoint = endpoint, kind = %mapped.kind(), detail = mapped.detail(), "WebSocket client connect failed" ); return Err(mapped); }, Err(_) => { let error = timeout_error("WebSocket client connect", config.connect_timeout()); tracing::warn!(target: TRACING_TARGET, endpoint = endpoint, timeout_ms = duration_millis(config.connect_timeout()), "WebSocket client connect timed out"); return Err(error); }, }; tracing::info!(target: TRACING_TARGET, endpoint = endpoint, "WebSocket client connected"); return Ok(crate::WebSocketConnection::from_client(stream, config)); } fn tungstenite_config(config: &crate::WebSocketConfig) -> tokio_tungstenite::tungstenite::protocol::WebSocketConfig { return tokio_tungstenite::tungstenite::protocol::WebSocketConfig::default() .write_buffer_size(config.write_buffer_size()) .max_write_buffer_size(config.max_write_buffer_size()) .max_message_size(Some(config.max_message_size())) .max_frame_size(Some(config.max_frame_size())); } fn timeout_error(operation: &str, timeout: std::time::Duration) -> game_realtime_transport_lib::TransportError { return game_realtime_transport_lib::TransportError::new( game_realtime_transport_lib::TransportErrorKind::Timeout, format!("{operation} exceeded configured deadline of {} ms", duration_millis(timeout)), ); } fn duration_millis(duration: std::time::Duration) -> u128 { return duration.as_millis(); } fn map_connect_error(error: tokio_tungstenite::tungstenite::Error) -> game_realtime_transport_lib::TransportError { let kind = match &error { tokio_tungstenite::tungstenite::Error::Url(_) => game_realtime_transport_lib::TransportErrorKind::InvalidConfiguration, _ => game_realtime_transport_lib::TransportErrorKind::Connect, }; return game_realtime_transport_lib::TransportError::new(kind, error.to_string()); } fn map_stream_error(error: tokio_tungstenite::tungstenite::Error) -> game_realtime_transport_lib::TransportError { let kind = match &error { tokio_tungstenite::tungstenite::Error::ConnectionClosed | tokio_tungstenite::tungstenite::Error::AlreadyClosed => { game_realtime_transport_lib::TransportErrorKind::Closed }, tokio_tungstenite::tungstenite::Error::Io(_) => game_realtime_transport_lib::TransportErrorKind::Io, tokio_tungstenite::tungstenite::Error::Capacity(_) => game_realtime_transport_lib::TransportErrorKind::MessageTooLarge, tokio_tungstenite::tungstenite::Error::WriteBufferFull(_) => game_realtime_transport_lib::TransportErrorKind::Backpressure, _ => game_realtime_transport_lib::TransportErrorKind::Protocol, }; return game_realtime_transport_lib::TransportError::new(kind, error.to_string()); } #[cfg(test)] #[path = "../unit_tests/websocket.rs"] mod tests;