use std::{net::SocketAddr, sync::Arc, time::Duration}; use tokio::net::{TcpListener, TcpStream, ToSocketAddrs}; use tokio_tungstenite::tungstenite::protocol::WebSocketConfig; use crate::{Error, session::Session}; /// Limits applied to every accepted connection. #[derive(Debug, Clone)] pub struct ServerConfig { /// Largest message (after reassembling fragments) a peer may send. pub max_message_size: usize, /// Largest single frame a peer may send. pub max_frame_size: usize, /// How long a client gets to complete the WebSocket handshake. pub handshake_timeout: Duration, } impl Default for ServerConfig { fn default() -> Self { Self { max_message_size: 1 << 20, max_frame_size: 1 << 20, handshake_timeout: Duration::from_secs(5), } } } /// A standalone WebSocket server that hands each connection to a callback as a /// [`Session`]. pub struct SessionServer { listener: TcpListener, config: ServerConfig, } impl SessionServer { pub async fn bind(addr: impl ToSocketAddrs) -> crate::Result { Ok(Self::from_listener(TcpListener::bind(addr).await?)) } pub fn from_listener(listener: TcpListener) -> Self { Self { listener, config: ServerConfig::default(), } } pub fn with_config(mut self, config: ServerConfig) -> Self { self.config = config; self } pub fn local_addr(&self) -> std::io::Result { self.listener.local_addr() } /// Accept one connection. The receiver is not started: register handlers, /// then call [`Session::start_receiver`]. pub async fn accept(&self) -> crate::Result<(Session, SocketAddr)> { let (stream, addr) = self.listener.accept().await?; Ok((handshake(stream, &self.config).await?, addr)) } /// Accept connections forever, running `on_conn` for each one. /// /// `on_conn` should register handlers and return; the receiver starts once /// it does, so no message is processed before its handler exists. If it /// returns an error the connection is closed. pub async fn session_loop(&self, on_conn: F) -> crate::Result<()> where F: Fn(Session, SocketAddr) -> Fut + Send + Sync + 'static, Fut: Future> + Send + 'static, { let on_conn = Arc::new(on_conn); loop { let (stream, addr) = match self.listener.accept().await { Ok(conn) => conn, Err(e) => { // e.g. out of file descriptors: back off instead of exiting. eprintln!("Accept failed: {e}"); tokio::time::sleep(Duration::from_millis(100)).await; continue; } }; let on_conn = on_conn.clone(); let config = self.config.clone(); tokio::spawn(async move { let session = match handshake(stream, &config).await { Ok(session) => session, Err(e) => { eprintln!("Handshake failed from {addr}: {e}"); return; } }; if let Err(e) = on_conn(session.clone(), addr).await { eprintln!("Connection error from {addr}: {e}"); let _ = session.close().await; return; } session.start_receiver(); }); } } } async fn handshake(stream: TcpStream, config: &ServerConfig) -> crate::Result { let ws_config = WebSocketConfig::default() .max_message_size(Some(config.max_message_size)) .max_frame_size(Some(config.max_frame_size)); let ws = tokio::time::timeout( config.handshake_timeout, tokio_tungstenite::accept_async_with_config(stream, Some(ws_config)), ) .await? .map_err(|e| Error::Transport(Box::new(e)))?; Ok(Session::from_tungstenite(ws)) }