#![cfg(all(feature = "server", feature = "client"))] use std::net::SocketAddr; use std::sync::Arc; use std::sync::atomic::{AtomicUsize, Ordering}; use std::time::Duration; use futures_util::{SinkExt, StreamExt}; use serde::{Deserialize, Serialize}; use session_rs::server::SessionServer; use session_rs::{Error, Method, Session}; use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::net::TcpStream; use tokio_tungstenite::tungstenite::Message as WsMessage; macro_rules! method { ($ty:ident, $name:literal, $req:ty, $res:ty) => { #[derive(Debug, Serialize, Deserialize)] struct $ty; impl Method for $ty { const NAME: &'static str = $name; type Request = $req; type Response = $res; type Error = String; } }; } method!(Echo, "echo", String, String); method!(Fail, "fail", String, String); method!(Panic, "panic", (), ()); method!(Numbers, "numbers", Vec, u32); method!(AskBack, "ask_back", String, String); method!(Notice, "notice", String, ()); method!(Silent, "silent", (), ()); const WAIT: Duration = Duration::from_secs(5); async fn register_handlers(session: &Session) { session.on_request::(async |_, s| Ok(s)).await; session.on_request::(async |_, s| Err(format!("failed: {s}"))).await; session .on_request::(async |_, ()| -> Result<(), String> { panic!("boom") }) .await; session .on_request::(async |_, v| Ok(v.iter().sum())) .await; session .on_request::({ let session = session.clone(); move |_, s| { let session = session.clone(); async move { // A handler requesting from its own peer must not deadlock. let reply = session.request::(format!("back:{s}")).await; reply.map_err(|e| e.to_string())? } } }) .await; session .on_request::(async |_, ()| { tokio::time::sleep(Duration::from_secs(60)).await; Ok(()) }) .await; } async fn start_server() -> SocketAddr { start_server_with(|_| {}).await } /// Starts a server; `on_session` sees every accepted session after its /// handlers are registered. async fn start_server_with(on_session: impl Fn(Session) + Send + Sync + 'static) -> SocketAddr { let server = SessionServer::bind("127.0.0.1:0").await.unwrap(); let addr = server.local_addr().unwrap(); let on_session = Arc::new(on_session); tokio::spawn(async move { server .session_loop(move |session, _| { let on_session = on_session.clone(); async move { // Registering late must not lose a request sent right after connecting. tokio::time::sleep(Duration::from_millis(50)).await; register_handlers(&session).await; on_session(session); Ok(()) } }) .await }); addr } async fn connect(addr: SocketAddr) -> Session { let session = Session::connect(&format!("ws://{addr}")).await.unwrap(); register_handlers(&session).await; session.start_receiver(); session } async fn raw_client(addr: SocketAddr) -> tokio_tungstenite::WebSocketStream> { tokio_tungstenite::connect_async(format!("ws://{addr}")).await.unwrap().0 } async fn next_text( ws: &mut tokio_tungstenite::WebSocketStream>, ) -> serde_json::Value { loop { match tokio::time::timeout(WAIT, ws.next()).await.unwrap().unwrap().unwrap() { WsMessage::Text(t) => return serde_json::from_str(&t).unwrap(), _ => continue, } } } #[tokio::test] async fn request_response_and_error() { let client = connect(start_server().await).await; assert_eq!(client.request::("hi".into()).await.unwrap(), Ok("hi".into())); assert_eq!(client.request::(vec![1, 2, 3]).await.unwrap(), Ok(6)); assert_eq!( client.request::("x".into()).await.unwrap(), Err("failed: x".into()) ); } #[tokio::test] async fn concurrent_requests_are_matched_by_id() { let client = connect(start_server().await).await; let replies = futures_util::future::join_all( (0..200).map(|i| { let client = client.clone(); async move { client.request::(i.to_string()).await.unwrap() } }), ) .await; for (i, reply) in replies.into_iter().enumerate() { assert_eq!(reply, Ok(i.to_string())); } } #[tokio::test] async fn wire_format_matches_protocol() { let mut ws = raw_client(start_server().await).await; ws.send(WsMessage::Text( r#"{"type":"request","id":7,"method":"echo","data":"hi"}"#.into(), )) .await .unwrap(); assert_eq!( next_text(&mut ws).await, serde_json::json!({"type": "response", "id": 7, "result": "hi"}) ); ws.send(WsMessage::Text( r#"{"type":"request","id":8,"method":"fail","data":"x"}"#.into(), )) .await .unwrap(); assert_eq!( next_text(&mut ws).await, serde_json::json!({"type": "errorresponse", "id": 8, "error": "failed: x"}) ); } #[tokio::test] async fn unknown_method_and_bad_data_get_error_responses() { let mut ws = raw_client(start_server().await).await; ws.send(WsMessage::Text( r#"{"type":"request","id":1,"method":"nope","data":null}"#.into(), )) .await .unwrap(); assert_eq!( next_text(&mut ws).await, serde_json::json!({"type": "errorresponse", "id": 1, "error": "Unknown method: nope"}) ); ws.send(WsMessage::Text( r#"{"type":"request","id":2,"method":"numbers","data":"not a list"}"#.into(), )) .await .unwrap(); let reply = next_text(&mut ws).await; assert_eq!(reply["type"], "errorresponse"); assert_eq!(reply["id"], 2); assert!(reply["error"].as_str().unwrap().starts_with("Invalid request data")); // Garbage text is ignored, not fatal. ws.send(WsMessage::Text("not json".into())).await.unwrap(); ws.send(WsMessage::Text( r#"{"type":"request","id":3,"method":"echo","data":"still here"}"#.into(), )) .await .unwrap(); assert_eq!(next_text(&mut ws).await["result"], "still here"); } #[tokio::test] async fn panicking_handler_fails_only_that_request() { let client = connect(start_server().await).await; assert_eq!( client.request::(()).await.unwrap(), Err("Handler panicked".into()) ); assert_eq!(client.request::("ok".into()).await.unwrap(), Ok("ok".into())); } #[tokio::test] async fn handler_can_request_from_its_peer() { let client = connect(start_server().await).await; assert_eq!( client.request::("x".into()).await.unwrap(), Ok("back:x".into()) ); } #[tokio::test] async fn notifications_reach_the_peer() { let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel(); let addr = start_server_with(move |session| { let session = session.clone(); tokio::spawn(async move { session.notify::("hello".into()).await.unwrap(); }); }) .await; let client = Session::connect(&format!("ws://{addr}")).await.unwrap(); client .on_notification::(move |msg| { let tx = tx.clone(); async move { tx.send(msg).unwrap(); } }) .await; client.start_receiver(); assert_eq!( tokio::time::timeout(WAIT, rx.recv()).await.unwrap(), Some("hello".into()) ); } #[tokio::test] async fn close_fails_pending_requests_and_runs_on_close_once() { let closes = Arc::new(AtomicUsize::new(0)); let (session_tx, mut session_rx) = tokio::sync::mpsc::unbounded_channel(); let addr = start_server_with(move |session| { session_tx.send(session).unwrap(); }) .await; let client = connect(addr).await; client .on_close({ let closes = closes.clone(); move || { closes.fetch_add(1, Ordering::SeqCst); async { Ok(()) } } }) .await; let pending = tokio::spawn({ let client = client.clone(); async move { client.request::(()).await } }); let server_side = tokio::time::timeout(WAIT, session_rx.recv()).await.unwrap().unwrap(); tokio::time::sleep(Duration::from_millis(50)).await; server_side.close().await.unwrap(); let result = tokio::time::timeout(WAIT, pending).await.unwrap().unwrap(); assert!(matches!(result, Err(Error::ConnectionClosed)), "{result:?}"); tokio::time::timeout(WAIT, client.closed()).await.unwrap(); client.close().await.unwrap(); assert_eq!(closes.load(Ordering::SeqCst), 1); assert!(matches!( client.request::("late".into()).await, Err(Error::ConnectionClosed) )); } #[tokio::test] async fn request_timeout() { let client = connect(start_server().await).await; let result = client .request_timeout::((), Duration::from_millis(100)) .await; assert!(matches!(result, Err(Error::Timeout)), "{result:?}"); } #[tokio::test] async fn oversized_message_closes_only_that_connection() { let addr = start_server().await; let mut ws = raw_client(addr).await; let _ = ws.send(WsMessage::Text("x".repeat(2 << 20).into())).await; let closed = tokio::time::timeout(WAIT, async { loop { match ws.next().await { None | Some(Err(_)) | Some(Ok(WsMessage::Close(_))) => break, _ => {} } } }) .await; assert!(closed.is_ok(), "server kept the connection open"); let client = connect(addr).await; assert_eq!(client.request::("alive".into()).await.unwrap(), Ok("alive".into())); } /// Opens a TCP connection and completes a WebSocket handshake by hand. async fn raw_handshake(addr: SocketAddr) -> TcpStream { let mut tcp = TcpStream::connect(addr).await.unwrap(); tcp.write_all( format!( "GET / HTTP/1.1\r\nHost: {addr}\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\ Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\nSec-WebSocket-Version: 13\r\n\r\n" ) .as_bytes(), ) .await .unwrap(); let mut response = Vec::new(); while !response.ends_with(b"\r\n\r\n") { response.push(tcp.read_u8().await.unwrap()); } assert!(response.starts_with(b"HTTP/1.1 101")); tcp } #[tokio::test] async fn huge_frame_length_header_is_rejected_without_allocating() { let addr = start_server().await; let mut tcp = raw_handshake(addr).await; // Masked text frame claiming a 2^62-byte payload. session-rs 0.1 tried to // allocate this up front. let mut frame = vec![0x81, 0x80 | 127]; frame.extend_from_slice(&(1u64 << 62).to_be_bytes()); frame.extend_from_slice(&[1, 2, 3, 4]); tcp.write_all(&frame).await.unwrap(); let mut buf = [0u8; 1024]; let closed = tokio::time::timeout(WAIT, async { loop { match tcp.read(&mut buf).await { Ok(0) | Err(_) => break, Ok(_) => {} } } }) .await; assert!(closed.is_ok(), "server kept the connection open"); let client = connect(addr).await; assert_eq!(client.request::("alive".into()).await.unwrap(), Ok("alive".into())); } #[tokio::test] async fn ping_timeout_closes_unresponsive_peer() { let closes = Arc::new(AtomicUsize::new(0)); let addr = start_server_with({ let closes = closes.clone(); move |session| { let closes = closes.clone(); tokio::spawn(async move { session .on_close(move || { closes.fetch_add(1, Ordering::SeqCst); async { Ok(()) } }) .await; session.start_ping(Duration::from_millis(50), Duration::from_millis(100)); }); } }) .await; // A tungstenite client answers pings, so it must stay connected. let client = connect(addr).await; tokio::time::sleep(Duration::from_millis(400)).await; assert_eq!(client.request::("alive".into()).await.unwrap(), Ok("alive".into())); assert_eq!(closes.load(Ordering::SeqCst), 0); // A raw TCP peer never reads, so it never pongs. let _tcp = raw_handshake(addr).await; tokio::time::timeout(WAIT, async { while closes.load(Ordering::SeqCst) == 0 { tokio::time::sleep(Duration::from_millis(20)).await; } }) .await .expect("unresponsive peer was not closed"); } #[cfg(feature = "axum")] #[tokio::test] async fn axum_adapter_serves_sessions_next_to_http_routes() { use axum::{Router, extract::WebSocketUpgrade, routing::get}; let app = Router::new() .route( "/", get(async |upgrade: WebSocketUpgrade| { upgrade.on_upgrade(async |socket| { let session = Session::from_axum(socket); register_handlers(&session).await; session.start_receiver(); }) }), ) .route("/health", get(async || "ok")); let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); let addr = listener.local_addr().unwrap(); tokio::spawn(async move { axum::serve(listener, app).await }); let client = connect(addr).await; assert_eq!(client.request::("hi".into()).await.unwrap(), Ok("hi".into())); assert_eq!( client.request::("y".into()).await.unwrap(), Ok("back:y".into()) ); let mut tcp = TcpStream::connect(addr).await.unwrap(); tcp.write_all(b"GET /health HTTP/1.1\r\nHost: x\r\nConnection: close\r\n\r\n") .await .unwrap(); let mut body = String::new(); tcp.read_to_string(&mut body).await.unwrap(); assert!(body.starts_with("HTTP/1.1 200") && body.ends_with("ok"), "{body}"); } #[tokio::test] async fn closing_releases_handlers_that_capture_the_session() { let client = connect(start_server().await).await; let captured = Arc::new(()); // Handlers typically hold a clone of their own session; that must not keep // the session (and everything the handlers capture) alive after close. client .on_request::({ let session = client.clone(); let captured = captured.clone(); move |_, ()| { let _ = (&session, &captured); async { Ok(()) } } }) .await; client .on_close({ let session = client.clone(); let captured = captured.clone(); move || { let _ = (&session, &captured); async { Ok(()) } } }) .await; assert_eq!(Arc::strong_count(&captured), 3); client.close().await.unwrap(); tokio::time::timeout(WAIT, async { while Arc::strong_count(&captured) > 1 { tokio::time::sleep(Duration::from_millis(10)).await; } }) .await .expect("handlers were not released after close"); }