diff --git a/src/lib.rs b/src/lib.rs index dc2e07e..ffa7420 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -6,6 +6,7 @@ mod error; mod handler; #[cfg(feature = "mocking")] mod mock_server; +mod socket_loop; #[cfg(feature = "msgpack")] pub use codec::MsgPackCodec; @@ -18,30 +19,17 @@ pub use mock_server::msgpack_echo_server; #[cfg(feature = "mocking")] pub use mock_server::{EchoControlMessage, auth_echo_server, echo_server, get_mock_address}; -use bytes::Bytes; -use futures::{SinkExt, StreamExt, stream::SplitSink, stream::SplitStream}; -use std::time::Duration; -use tokio::{ - net::TcpStream, - select, - sync::{mpsc, oneshot}, - time::sleep, -}; -use tokio_tungstenite::{ - MaybeTlsStream, WebSocketStream, connect_async, - tungstenite::{self, Message, Utf8Bytes, protocol::CloseFrame}, -}; +pub(crate) use socket_loop::WebSocketStreamType; +use socket_loop::{TxChannelPayload, send_close, socket_loop_split}; + +use futures::StreamExt; +use tokio::sync::{mpsc, oneshot}; +use tokio_tungstenite::{connect_async, tungstenite::Message}; #[cfg(feature = "tracing")] -use tracing::{debug, error, info, instrument, trace}; +use tracing::{debug, info, instrument, trace}; use url::Url; -#[derive(Debug)] -struct TxChannelPayload { - message: Message, - response_tx: oneshot::Sender>, -} - /// A WebSocket client that manages the connection to a WebSocket server. /// The client can send and receive messages, and will transparently handle protocol messages. /// @@ -281,750 +269,3 @@ where } } } - -pub(crate) type WebSocketStreamType = WebSocketStream>; -type SocketSink = SplitSink; -type SocketStream = SplitStream; - -enum LoopState { - Running, - Error(Error), - Closed, -} - -/// Send a close frame via the tx channel and wait for confirmation. -async fn send_close(sender: &mpsc::Sender) -> Result<(), Error> { - let (tx, rx) = oneshot::channel::>(); - sender - .send(TxChannelPayload { - message: Message::Close(Some(CloseFrame { - code: tungstenite::protocol::frame::coding::CloseCode::Normal, - reason: Utf8Bytes::from_static("Closing Connection"), - })), - response_tx: tx, - }) - .await - .map_err(|_| Error::WebsocketClosed)?; - match rx.await { - Ok(result) => result, - Err(_) => unreachable!("Socket loop always sends response before dropping one-shot"), - } -} - -#[cfg_attr( - feature = "tracing", - instrument(skip(keepalive_interval, keepalive_message)) -)] -async fn socket_loop_split( - mut receiver: mpsc::Receiver, - mut sender: mpsc::Sender, - mut sink: SocketSink, - mut stream: SocketStream, - keepalive_interval: Option, - keepalive_message: Option, -) -> Result<(), Error> { - let mut state = LoopState::Running; - while matches!(state, LoopState::Running) { - state = if let Some(interval) = keepalive_interval { - select! { - outgoing_message = receiver.recv() => send_socket_message(outgoing_message, &mut sink).await, - incoming_message = stream.next() => socket_message_received(incoming_message, &mut sender, &mut sink).await, - () = sleep(interval) => send_keepalive(&mut sink, keepalive_message.as_ref()).await, - } - } else { - select! { - outgoing_message = receiver.recv() => send_socket_message(outgoing_message, &mut sink).await, - incoming_message = stream.next() => socket_message_received(incoming_message, &mut sender, &mut sink).await, - } - }; - } - match state { - LoopState::Error(e) => Err(e), - LoopState::Closed => Ok(()), - LoopState::Running => unreachable!("We only exit when closed or errored"), - } -} - -#[cfg_attr(feature = "tracing", instrument)] -async fn send_socket_message( - message: Option, - sink: &mut SocketSink, -) -> LoopState { - if let Some(message) = message { - #[cfg(feature = "tracing")] - trace!("Sending message: {:?}", message); - let send_result = sink.send(message.message).await.map_err(Error::from); - let socket_error = send_result.is_err(); - match message.response_tx.send(send_result) { - Ok(()) => { - if socket_error { - LoopState::Error(Error::WebsocketClosed) - } else { - LoopState::Running - } - } - Err(_) => LoopState::Error(Error::SocketeerDroppedWithoutClosing), - } - } else { - #[cfg(feature = "tracing")] - error!("Socketeer dropped without closing connection"); - LoopState::Error(Error::SocketeerDroppedWithoutClosing) - } -} - -#[cfg_attr(feature = "tracing", instrument)] -async fn socket_message_received( - message: Option>, - sender: &mut mpsc::Sender, - sink: &mut SocketSink, -) -> LoopState { - const PONG_BYTES: Bytes = Bytes::from_static(b"pong"); - match message { - Some(Ok(message)) => match message { - Message::Ping(_) => { - let send_result = sink - .send(Message::Pong(PONG_BYTES)) - .await - .map_err(Error::from); - match send_result { - Ok(()) => LoopState::Running, - Err(e) => { - #[cfg(feature = "tracing")] - error!("Error sending Pong: {:?}", e); - LoopState::Error(e) - } - } - } - Message::Close(_) => { - let close_result = sink.close().await; - match close_result { - Ok(()) => LoopState::Closed, - Err(e) => { - #[cfg(feature = "tracing")] - error!("Error sending Close: {:?}", e); - LoopState::Error(Error::from(e)) - } - } - } - Message::Text(_) | Message::Binary(_) => match sender.send(message).await { - Ok(()) => LoopState::Running, - Err(_) => LoopState::Error(Error::SocketeerDroppedWithoutClosing), - }, - _ => LoopState::Running, - }, - Some(Err(e)) => { - #[cfg(feature = "tracing")] - error!("Error receiving message: {:?}", e); - LoopState::Error(Error::WebsocketError(e)) - } - None => { - #[cfg(feature = "tracing")] - info!("Websocket Closed, closing rx channel"); - LoopState::Error(Error::WebsocketClosed) - } - } -} - -#[cfg_attr(feature = "tracing", instrument)] -async fn send_keepalive(sink: &mut SocketSink, custom_message: Option<&Message>) -> LoopState { - let message = if let Some(custom) = custom_message { - #[cfg(feature = "tracing")] - trace!("Timeout waiting for message, sending custom keepalive"); - custom.clone() - } else { - #[cfg(feature = "tracing")] - trace!("Timeout waiting for message, sending Ping"); - Message::Ping(Bytes::new()) - }; - let result = sink.send(message).await.map_err(Error::from); - match result { - Ok(()) => LoopState::Running, - Err(e) => { - #[cfg(feature = "tracing")] - error!("Error sending keepalive: {:?}", e); - LoopState::Error(e) - } - } -} - -#[cfg(all(test, feature = "mocking"))] -mod tests { - use super::*; - use tokio::time::sleep; - - type EchoJson = JsonCodec; - - #[tokio::test] - async fn test_server_startup() { - let _server_address = get_mock_address(echo_server).await; - } - - #[tokio::test] - async fn test_connection() { - let server_address = get_mock_address(echo_server).await; - let _socketeer: Socketeer = Socketeer::connect(&format!("ws://{server_address}")) - .await - .unwrap(); - } - - #[tokio::test] - async fn test_bad_url() { - let error: Result, Error> = Socketeer::connect("Not a URL").await; - assert!(matches!(error.unwrap_err(), Error::UrlParse { .. })); - } - - #[tokio::test] - async fn test_send_receive() { - let server_address = get_mock_address(echo_server).await; - let mut socketeer: Socketeer = - Socketeer::connect(&format!("ws://{server_address}")) - .await - .unwrap(); - let message = EchoControlMessage::Message("Hello".to_string()); - socketeer.send(message.clone()).await.unwrap(); - let received_message = socketeer.next_message().await.unwrap(); - assert_eq!(message, received_message); - } - - #[tokio::test] - async fn test_ping_request() { - let server_address = get_mock_address(echo_server).await; - let mut socketeer: Socketeer = - Socketeer::connect(&format!("ws://{server_address}")) - .await - .unwrap(); - let ping_request = EchoControlMessage::SendPing; - socketeer.send(ping_request).await.unwrap(); - // The server will respond with a ping request, which Socketeer will transparently respond to - let message = EchoControlMessage::Message("Hello".to_string()); - socketeer.send(message.clone()).await.unwrap(); - let received_message = socketeer.next_message().await.unwrap(); - assert_eq!(received_message, message); - // We should send a ping in here - sleep(Duration::from_millis(2200)).await; - // Ensure everything shuts down so we exercize the ping functionality fully - socketeer.close_connection().await.unwrap(); - } - - #[tokio::test] - async fn test_reconnection() { - let server_address = get_mock_address(echo_server).await; - let mut socketeer: Socketeer = - Socketeer::connect(&format!("ws://{server_address}")) - .await - .unwrap(); - let message = EchoControlMessage::Message("Hello".to_string()); - socketeer.send(message.clone()).await.unwrap(); - let received_message = socketeer.next_message().await.unwrap(); - assert_eq!(message, received_message); - socketeer = socketeer.reconnect().await.unwrap(); - let message = EchoControlMessage::Message("Hello".to_string()); - socketeer.send(message.clone()).await.unwrap(); - let received_message = socketeer.next_message().await.unwrap(); - assert_eq!(message, received_message); - socketeer.close_connection().await.unwrap(); - } - - #[tokio::test] - async fn test_closed_socket() { - let server_address = get_mock_address(echo_server).await; - let mut socketeer: Socketeer = - Socketeer::connect(&format!("ws://{server_address}")) - .await - .unwrap(); - let close_request = EchoControlMessage::Close; - socketeer.send(close_request.clone()).await.unwrap(); - let response = socketeer.next_message().await; - assert!(matches!(response.unwrap_err(), Error::WebsocketClosed)); - let send_result = socketeer.send(close_request).await; - assert!(send_result.is_err()); - let error = send_result.unwrap_err(); - println!("Actual Error: {error:#?}"); - assert!(matches!(error, Error::WebsocketClosed)); - } - - #[tokio::test] - async fn test_close_request() { - let server_address = get_mock_address(echo_server).await; - let socketeer: Socketeer = Socketeer::connect(&format!("ws://{server_address}")) - .await - .unwrap(); - socketeer.close_connection().await.unwrap(); - } - - #[tokio::test] - async fn test_connect_with_default_options() { - let server_address = get_mock_address(echo_server).await; - let mut socketeer: Socketeer = - Socketeer::connect_with(&format!("ws://{server_address}"), ConnectOptions::default()) - .await - .unwrap(); - let message = EchoControlMessage::Message("Hello".to_string()); - socketeer.send(message.clone()).await.unwrap(); - let received_message = socketeer.next_message().await.unwrap(); - assert_eq!(message, received_message); - } - - #[tokio::test] - async fn test_raw_codec_message_roundtrip() { - // Typed send/next_message round-trip when the codec is RawCodec — the - // codec is identity, so frames pass through unchanged. - let server_address = get_mock_address(echo_server).await; - let mut socketeer: Socketeer = - Socketeer::connect(&format!("ws://{server_address}")) - .await - .unwrap(); - let raw_text = r#"{"Message":"raw hello"}"#; - socketeer - .send(Message::Text(raw_text.into())) - .await - .unwrap(); - let received = socketeer.next_message().await.unwrap(); - assert_eq!(received, Message::Text(raw_text.into())); - } - - #[tokio::test] - async fn test_disabled_keepalive() { - let server_address = get_mock_address(echo_server).await; - let options = ConnectOptions { - keepalive_interval: None, - ..ConnectOptions::default() - }; - let mut socketeer: Socketeer = - Socketeer::connect_with(&format!("ws://{server_address}"), options) - .await - .unwrap(); - let message = EchoControlMessage::Message("Hello".to_string()); - socketeer.send(message.clone()).await.unwrap(); - let received_message = socketeer.next_message().await.unwrap(); - assert_eq!(message, received_message); - } - - #[tokio::test] - async fn test_handler_on_connected() { - use serde::{Deserialize, Serialize}; - use std::sync::Arc; - use tokio::sync::Mutex; - - #[derive(Clone, Debug, Serialize, Deserialize, PartialEq)] - struct AuthResponse { - status: String, - } - - struct TestAuthHandler { - connected_count: Arc>, - } - - impl ConnectionHandler for TestAuthHandler { - async fn on_connected( - &mut self, - ctx: &mut HandshakeContext<'_, C>, - ) -> Result<(), Error> { - ctx.send_text(r#"{"action":"auth","token":"test-token"}"#) - .await?; - let text = ctx.recv_text().await?; - let response: AuthResponse = serde_json::from_str(&text).unwrap(); - assert_eq!(response.status, "authenticated"); - let mut count = self.connected_count.lock().await; - *count += 1; - Ok(()) - } - } - - let connected_count = Arc::new(Mutex::new(0u32)); - let handler = TestAuthHandler { - connected_count: connected_count.clone(), - }; - - let server_address = get_mock_address(auth_echo_server).await; - let mut socketeer: Socketeer = Socketeer::connect_with_codec( - &format!("ws://{server_address}"), - ConnectOptions::default(), - JsonCodec::new(), - handler, - ) - .await - .unwrap(); - - assert_eq!(*connected_count.lock().await, 1); - - let message = EchoControlMessage::Message("after auth".to_string()); - socketeer.send(message.clone()).await.unwrap(); - let received = socketeer.next_message().await.unwrap(); - assert_eq!(message, received); - } - - #[tokio::test] - async fn test_handler_reconnect() { - use std::sync::Arc; - use tokio::sync::Mutex; - - struct ReconnectHandler { - connected_count: Arc>, - disconnected_count: Arc>, - } - - impl ConnectionHandler for ReconnectHandler { - async fn on_connected( - &mut self, - ctx: &mut HandshakeContext<'_, C>, - ) -> Result<(), Error> { - ctx.send_text(r#"{"action":"auth","token":"test-token"}"#) - .await?; - let _response = ctx.recv_text().await?; - let mut count = self.connected_count.lock().await; - *count += 1; - Ok(()) - } - - async fn on_disconnected(&mut self) { - let mut count = self.disconnected_count.lock().await; - *count += 1; - } - } - - let connected_count = Arc::new(Mutex::new(0u32)); - let disconnected_count = Arc::new(Mutex::new(0u32)); - let handler = ReconnectHandler { - connected_count: connected_count.clone(), - disconnected_count: disconnected_count.clone(), - }; - - let server_address = get_mock_address(auth_echo_server).await; - let mut socketeer = Socketeer::::connect_with_codec( - &format!("ws://{server_address}"), - ConnectOptions::default(), - JsonCodec::new(), - handler, - ) - .await - .unwrap(); - - assert_eq!(*connected_count.lock().await, 1); - assert_eq!(*disconnected_count.lock().await, 0); - - // Send a message to verify connection works - let message = EchoControlMessage::Message("before reconnect".to_string()); - socketeer.send(message.clone()).await.unwrap(); - let received = socketeer.next_message().await.unwrap(); - assert_eq!(message, received); - - // Reconnect — handler should fire again - socketeer = socketeer.reconnect().await.unwrap(); - - assert_eq!(*connected_count.lock().await, 2); - assert_eq!(*disconnected_count.lock().await, 1); - - // Verify connection still works after reconnect - let message = EchoControlMessage::Message("after reconnect".to_string()); - socketeer.send(message.clone()).await.unwrap(); - let received = socketeer.next_message().await.unwrap(); - assert_eq!(message, received); - - socketeer.close_connection().await.unwrap(); - } - - #[cfg(feature = "msgpack")] - #[tokio::test] - async fn test_msgpack_send_receive() { - type EchoMsgPack = MsgPackCodec; - - let server_address = get_mock_address(msgpack_echo_server).await; - let mut socketeer: Socketeer = - Socketeer::connect(&format!("ws://{server_address}")) - .await - .unwrap(); - let message = EchoControlMessage::Message("msgpack hello".to_string()); - socketeer.send(message.clone()).await.unwrap(); - let received = socketeer.next_message().await.unwrap(); - assert_eq!(message, received); - socketeer.close_connection().await.unwrap(); - } - - #[tokio::test] - async fn test_handler_uses_codec_driven_send_recv() { - // Exercises HandshakeContext::send / recv (the codec-driven path). - // Other handler tests only cover the raw send_text / recv_text helpers. - struct TypedHandshakeHandler; - - impl ConnectionHandler for TypedHandshakeHandler { - async fn on_connected( - &mut self, - ctx: &mut HandshakeContext<'_, EchoJson>, - ) -> Result<(), Error> { - ctx.send(&EchoControlMessage::Message("handshake".into())) - .await?; - let echoed = ctx.recv().await?; - assert_eq!(echoed, EchoControlMessage::Message("handshake".into())); - Ok(()) - } - } - - let server_address = get_mock_address(echo_server).await; - let mut socketeer: Socketeer = - Socketeer::connect_with_codec( - &format!("ws://{server_address}"), - ConnectOptions::default(), - JsonCodec::new(), - TypedHandshakeHandler, - ) - .await - .unwrap(); - - // Confirm normal traffic still flows after the typed handshake. - let message = EchoControlMessage::Message("after handshake".into()); - socketeer.send(message.clone()).await.unwrap(); - assert_eq!(socketeer.next_message().await.unwrap(), message); - socketeer.close_connection().await.unwrap(); - } - - #[tokio::test] - async fn test_handshake_recv_close_with_raw_codec() { - // Regression: with RawCodec, recv_raw returns Ok(Message::Close(_)) and - // RawCodec::decode is the identity, so a peer-initiated close used to - // surface as Ok(Close) instead of Err(WebsocketClosed). recv must - // intercept Close before delegating to the codec. - struct CloseExpecting; - - impl ConnectionHandler for CloseExpecting { - async fn on_connected( - &mut self, - ctx: &mut HandshakeContext<'_, RawCodec>, - ) -> Result<(), Error> { - // Ask the echo server to close (JSON unit-variant for EchoControlMessage::Close). - ctx.send(&Message::Text(r#""Close""#.into())).await?; - let err = ctx.recv().await.unwrap_err(); - assert!(matches!(err, Error::WebsocketClosed)); - Ok(()) - } - } - - let server_address = get_mock_address(echo_server).await; - let _socketeer: Socketeer = Socketeer::connect_with_codec( - &format!("ws://{server_address}"), - ConnectOptions::default(), - RawCodec::new(), - CloseExpecting, - ) - .await - .unwrap(); - } - - #[tokio::test] - async fn test_extra_headers_used() { - // Cover ConnectOptions::build_request's loop body that copies - // `extra_headers` onto the upgrade request. - let server_address = get_mock_address(echo_server).await; - let mut headers = tokio_tungstenite::tungstenite::http::HeaderMap::new(); - headers.insert("X-Test-Header", "socketeer".parse().unwrap()); - let options = ConnectOptions { - extra_headers: headers, - ..ConnectOptions::default() - }; - let mut socketeer: Socketeer = - Socketeer::connect_with(&format!("ws://{server_address}"), options) - .await - .unwrap(); - let message = EchoControlMessage::Message("hi".into()); - socketeer.send(message.clone()).await.unwrap(); - assert_eq!(socketeer.next_message().await.unwrap(), message); - socketeer.close_connection().await.unwrap(); - } - - #[tokio::test] - async fn test_auth_handler_bad_token() { - // Covers auth_echo_server's bad-token branch (sends {"status":"error"} - // and shuts down). The handler observes the error response, returns - // Ok, then a subsequent send fails because the server has closed. - struct BadTokenHandler; - - impl ConnectionHandler for BadTokenHandler { - async fn on_connected( - &mut self, - ctx: &mut HandshakeContext<'_, C>, - ) -> Result<(), Error> { - ctx.send_text(r#"{"action":"auth","token":"WRONG"}"#) - .await?; - let resp = ctx.recv_text().await?; - assert!(resp.contains("error")); - Ok(()) - } - } - - let server_address = get_mock_address(auth_echo_server).await; - let _socketeer: Socketeer = Socketeer::connect_with_codec( - &format!("ws://{server_address}"), - ConnectOptions::default(), - JsonCodec::new(), - BadTokenHandler, - ) - .await - .unwrap(); - } - - #[cfg(feature = "msgpack")] - #[tokio::test] - async fn test_msgpack_send_ping() { - // Covers the SendPing arm of msgpack_echo_server. - type EchoMsgPack = MsgPackCodec; - - let server_address = get_mock_address(msgpack_echo_server).await; - let mut socketeer: Socketeer = - Socketeer::connect(&format!("ws://{server_address}")) - .await - .unwrap(); - socketeer.send(EchoControlMessage::SendPing).await.unwrap(); - // Server replies with a Ping; Socketeer auto-Pongs. Round-trip a real - // message to confirm the connection is still alive. - let message = EchoControlMessage::Message("after ping".into()); - socketeer.send(message.clone()).await.unwrap(); - assert_eq!(socketeer.next_message().await.unwrap(), message); - socketeer.close_connection().await.unwrap(); - } - - #[cfg(feature = "msgpack")] - #[tokio::test] - async fn test_msgpack_close_request() { - // Covers the Close arm of msgpack_echo_server. - type EchoMsgPack = MsgPackCodec; - - let server_address = get_mock_address(msgpack_echo_server).await; - let mut socketeer: Socketeer = - Socketeer::connect(&format!("ws://{server_address}")) - .await - .unwrap(); - socketeer.send(EchoControlMessage::Close).await.unwrap(); - let result = socketeer.next_message().await; - assert!(matches!(result.unwrap_err(), Error::WebsocketClosed)); - } - - #[tokio::test] - async fn test_socketeer_debug_format() { - let server_address = get_mock_address(echo_server).await; - let socketeer: Socketeer = Socketeer::connect(&format!("ws://{server_address}")) - .await - .unwrap(); - let formatted = format!("{socketeer:?}"); - assert!(formatted.starts_with("Socketeer")); - assert!(formatted.contains("url")); - } - - #[tokio::test] - async fn test_send_raw_next_raw_message() { - // Cover the raw send/receive escape hatches on a typed (non-RawCodec) - // connection: send_raw bypasses encoding, next_raw_message bypasses - // decoding, so we can speak frames the codec wouldn't otherwise - // produce or accept. - let server_address = get_mock_address(echo_server).await; - let mut socketeer: Socketeer = - Socketeer::connect(&format!("ws://{server_address}")) - .await - .unwrap(); - let raw_text = r#"{"Message":"raw recv"}"#; - socketeer - .send_raw(Message::Text(raw_text.into())) - .await - .unwrap(); - let frame = socketeer.next_raw_message().await.unwrap(); - assert_eq!(frame, Message::Text(raw_text.into())); - socketeer.close_connection().await.unwrap(); - } - - #[cfg(feature = "msgpack")] - #[tokio::test] - async fn test_handshake_send_binary_recv_raw() { - // Cover HandshakeContext::send_binary by sending a pre-encoded - // msgpack frame from on_connected and reading the binary echo back - // via recv_raw. - struct BinaryHandshake; - - type EchoMsgPack = MsgPackCodec; - - impl ConnectionHandler for BinaryHandshake { - async fn on_connected( - &mut self, - ctx: &mut HandshakeContext<'_, EchoMsgPack>, - ) -> Result<(), Error> { - let payload = - rmp_serde::to_vec_named(&EchoControlMessage::Message("binary".into())).unwrap(); - ctx.send_binary(payload).await?; - let echo = ctx.recv_raw().await?; - assert!(matches!(echo, Message::Binary(_))); - Ok(()) - } - } - - let server_address = get_mock_address(msgpack_echo_server).await; - let socketeer: Socketeer = Socketeer::connect_with_codec( - &format!("ws://{server_address}"), - ConnectOptions::default(), - MsgPackCodec::new(), - BinaryHandshake, - ) - .await - .unwrap(); - socketeer.close_connection().await.unwrap(); - } - - #[cfg(feature = "msgpack")] - #[tokio::test] - async fn test_handshake_recv_text_rejects_binary() { - // Cover the non-Text branch of HandshakeContext::recv_text by pointing - // it at a server that only speaks binary frames. - struct ExpectsTextOnBinary; - - type EchoMsgPack = MsgPackCodec; - - impl ConnectionHandler for ExpectsTextOnBinary { - async fn on_connected( - &mut self, - ctx: &mut HandshakeContext<'_, EchoMsgPack>, - ) -> Result<(), Error> { - let payload = - rmp_serde::to_vec_named(&EchoControlMessage::Message("hi".into())).unwrap(); - ctx.send_binary(payload).await?; - // recv_text must reject the echoed Binary frame. - let err = ctx.recv_text().await.unwrap_err(); - assert!(matches!(err, Error::UnexpectedMessageType(_))); - Ok(()) - } - } - - let server_address = get_mock_address(msgpack_echo_server).await; - let socketeer: Socketeer = Socketeer::connect_with_codec( - &format!("ws://{server_address}"), - ConnectOptions::default(), - MsgPackCodec::new(), - ExpectsTextOnBinary, - ) - .await - .unwrap(); - socketeer.close_connection().await.unwrap(); - } - - #[tokio::test] - async fn test_binary_custom_keepalive() { - // The widening of custom_keepalive_message from Option to - // Option is otherwise unexercised. echo_server silently - // ignores Binary frames, so the receive queue stays clean and we can - // verify the connection survives a binary keepalive cycle. - let server_address = get_mock_address(echo_server).await; - let options = ConnectOptions { - keepalive_interval: Some(Duration::from_millis(100)), - custom_keepalive_message: Some(Message::Binary(Bytes::from_static(b"keepalive"))), - ..ConnectOptions::default() - }; - let mut socketeer: Socketeer = - Socketeer::connect_with(&format!("ws://{server_address}"), options) - .await - .unwrap(); - - // Wait long enough for at least a couple of keepalive ticks to fire. - sleep(Duration::from_millis(350)).await; - - let message = EchoControlMessage::Message("post-keepalive".into()); - socketeer.send(message.clone()).await.unwrap(); - assert_eq!(socketeer.next_message().await.unwrap(), message); - socketeer.close_connection().await.unwrap(); - } -} diff --git a/src/socket_loop.rs b/src/socket_loop.rs new file mode 100644 index 0000000..0e222fa --- /dev/null +++ b/src/socket_loop.rs @@ -0,0 +1,201 @@ +//! Internal background task that owns the split WebSocket stream and mediates +//! between the public [`crate::Socketeer`] handle and the network. +//! +//! The handle never touches the socket directly. It communicates with this loop +//! exclusively through the mpsc channels set up in +//! [`crate::Socketeer::connect_with_codec`]: outgoing frames arrive as +//! [`TxChannelPayload`] on the tx channel, and incoming data frames are pushed +//! onto the rx channel. Protocol frames (ping/pong/close) are handled here and +//! never surface to the consumer. + +use bytes::Bytes; +use futures::{SinkExt, StreamExt, stream::SplitSink, stream::SplitStream}; +use std::time::Duration; +use tokio::{ + net::TcpStream, + select, + sync::{mpsc, oneshot}, + time::sleep, +}; +use tokio_tungstenite::{ + MaybeTlsStream, WebSocketStream, + tungstenite::{self, Message, Utf8Bytes, protocol::CloseFrame}, +}; + +#[cfg(feature = "tracing")] +use tracing::{error, info, instrument, trace}; + +use crate::Error; + +/// A frame to send, paired with a one-shot the loop uses to report the result +/// of the underlying `sink.send`. +#[derive(Debug)] +pub(crate) struct TxChannelPayload { + pub(crate) message: Message, + pub(crate) response_tx: oneshot::Sender>, +} + +pub(crate) type WebSocketStreamType = WebSocketStream>; +type SocketSink = SplitSink; +type SocketStream = SplitStream; + +enum LoopState { + Running, + Error(Error), + Closed, +} + +/// Send a close frame via the tx channel and wait for confirmation. +pub(crate) async fn send_close(sender: &mpsc::Sender) -> Result<(), Error> { + let (tx, rx) = oneshot::channel::>(); + sender + .send(TxChannelPayload { + message: Message::Close(Some(CloseFrame { + code: tungstenite::protocol::frame::coding::CloseCode::Normal, + reason: Utf8Bytes::from_static("Closing Connection"), + })), + response_tx: tx, + }) + .await + .map_err(|_| Error::WebsocketClosed)?; + match rx.await { + Ok(result) => result, + Err(_) => unreachable!("Socket loop always sends response before dropping one-shot"), + } +} + +#[cfg_attr( + feature = "tracing", + instrument(skip(keepalive_interval, keepalive_message)) +)] +pub(crate) async fn socket_loop_split( + mut receiver: mpsc::Receiver, + mut sender: mpsc::Sender, + mut sink: SocketSink, + mut stream: SocketStream, + keepalive_interval: Option, + keepalive_message: Option, +) -> Result<(), Error> { + let mut state = LoopState::Running; + while matches!(state, LoopState::Running) { + state = if let Some(interval) = keepalive_interval { + select! { + outgoing_message = receiver.recv() => send_socket_message(outgoing_message, &mut sink).await, + incoming_message = stream.next() => socket_message_received(incoming_message, &mut sender, &mut sink).await, + () = sleep(interval) => send_keepalive(&mut sink, keepalive_message.as_ref()).await, + } + } else { + select! { + outgoing_message = receiver.recv() => send_socket_message(outgoing_message, &mut sink).await, + incoming_message = stream.next() => socket_message_received(incoming_message, &mut sender, &mut sink).await, + } + }; + } + match state { + LoopState::Error(e) => Err(e), + LoopState::Closed => Ok(()), + LoopState::Running => unreachable!("We only exit when closed or errored"), + } +} + +#[cfg_attr(feature = "tracing", instrument)] +async fn send_socket_message( + message: Option, + sink: &mut SocketSink, +) -> LoopState { + if let Some(message) = message { + #[cfg(feature = "tracing")] + trace!("Sending message: {:?}", message); + let send_result = sink.send(message.message).await.map_err(Error::from); + let socket_error = send_result.is_err(); + match message.response_tx.send(send_result) { + Ok(()) => { + if socket_error { + LoopState::Error(Error::WebsocketClosed) + } else { + LoopState::Running + } + } + Err(_) => LoopState::Error(Error::SocketeerDroppedWithoutClosing), + } + } else { + #[cfg(feature = "tracing")] + error!("Socketeer dropped without closing connection"); + LoopState::Error(Error::SocketeerDroppedWithoutClosing) + } +} + +#[cfg_attr(feature = "tracing", instrument)] +async fn socket_message_received( + message: Option>, + sender: &mut mpsc::Sender, + sink: &mut SocketSink, +) -> LoopState { + const PONG_BYTES: Bytes = Bytes::from_static(b"pong"); + match message { + Some(Ok(message)) => match message { + Message::Ping(_) => { + let send_result = sink + .send(Message::Pong(PONG_BYTES)) + .await + .map_err(Error::from); + match send_result { + Ok(()) => LoopState::Running, + Err(e) => { + #[cfg(feature = "tracing")] + error!("Error sending Pong: {:?}", e); + LoopState::Error(e) + } + } + } + Message::Close(_) => { + let close_result = sink.close().await; + match close_result { + Ok(()) => LoopState::Closed, + Err(e) => { + #[cfg(feature = "tracing")] + error!("Error sending Close: {:?}", e); + LoopState::Error(Error::from(e)) + } + } + } + Message::Text(_) | Message::Binary(_) => match sender.send(message).await { + Ok(()) => LoopState::Running, + Err(_) => LoopState::Error(Error::SocketeerDroppedWithoutClosing), + }, + _ => LoopState::Running, + }, + Some(Err(e)) => { + #[cfg(feature = "tracing")] + error!("Error receiving message: {:?}", e); + LoopState::Error(Error::WebsocketError(e)) + } + None => { + #[cfg(feature = "tracing")] + info!("Websocket Closed, closing rx channel"); + LoopState::Error(Error::WebsocketClosed) + } + } +} + +#[cfg_attr(feature = "tracing", instrument)] +async fn send_keepalive(sink: &mut SocketSink, custom_message: Option<&Message>) -> LoopState { + let message = if let Some(custom) = custom_message { + #[cfg(feature = "tracing")] + trace!("Timeout waiting for message, sending custom keepalive"); + custom.clone() + } else { + #[cfg(feature = "tracing")] + trace!("Timeout waiting for message, sending Ping"); + Message::Ping(Bytes::new()) + }; + let result = sink.send(message).await.map_err(Error::from); + match result { + Ok(()) => LoopState::Running, + Err(e) => { + #[cfg(feature = "tracing")] + error!("Error sending keepalive: {:?}", e); + LoopState::Error(e) + } + } +} diff --git a/tests/integration.rs b/tests/integration.rs new file mode 100644 index 0000000..1b043f5 --- /dev/null +++ b/tests/integration.rs @@ -0,0 +1,578 @@ +//! End-to-end tests that drive a real `Socketeer` against the in-crate mock +//! servers over loopback TCP. These exercise only the public API, so they live +//! here as integration tests rather than inline unit tests. +#![cfg(feature = "mocking")] + +use std::time::Duration; + +use bytes::Bytes; +use tokio::time::sleep; +use tokio_tungstenite::tungstenite::Message; + +use socketeer::{ + Codec, ConnectOptions, ConnectionHandler, EchoControlMessage, Error, HandshakeContext, + JsonCodec, RawCodec, Socketeer, auth_echo_server, echo_server, get_mock_address, +}; + +#[cfg(feature = "msgpack")] +use socketeer::{MsgPackCodec, msgpack_echo_server}; + +type EchoJson = JsonCodec; + +#[tokio::test] +async fn test_server_startup() { + let _server_address = get_mock_address(echo_server).await; +} + +#[tokio::test] +async fn test_connection() { + let server_address = get_mock_address(echo_server).await; + let _socketeer: Socketeer = Socketeer::connect(&format!("ws://{server_address}")) + .await + .unwrap(); +} + +#[tokio::test] +async fn test_bad_url() { + let error: Result, Error> = Socketeer::connect("Not a URL").await; + assert!(matches!(error.unwrap_err(), Error::UrlParse { .. })); +} + +#[tokio::test] +async fn test_send_receive() { + let server_address = get_mock_address(echo_server).await; + let mut socketeer: Socketeer = Socketeer::connect(&format!("ws://{server_address}")) + .await + .unwrap(); + let message = EchoControlMessage::Message("Hello".to_string()); + socketeer.send(message.clone()).await.unwrap(); + let received_message = socketeer.next_message().await.unwrap(); + assert_eq!(message, received_message); +} + +#[tokio::test] +async fn test_ping_request() { + let server_address = get_mock_address(echo_server).await; + let mut socketeer: Socketeer = Socketeer::connect(&format!("ws://{server_address}")) + .await + .unwrap(); + let ping_request = EchoControlMessage::SendPing; + socketeer.send(ping_request).await.unwrap(); + // The server will respond with a ping request, which Socketeer will transparently respond to + let message = EchoControlMessage::Message("Hello".to_string()); + socketeer.send(message.clone()).await.unwrap(); + let received_message = socketeer.next_message().await.unwrap(); + assert_eq!(received_message, message); + // We should send a ping in here + sleep(Duration::from_millis(2200)).await; + // Ensure everything shuts down so we exercize the ping functionality fully + socketeer.close_connection().await.unwrap(); +} + +#[tokio::test] +async fn test_reconnection() { + let server_address = get_mock_address(echo_server).await; + let mut socketeer: Socketeer = Socketeer::connect(&format!("ws://{server_address}")) + .await + .unwrap(); + let message = EchoControlMessage::Message("Hello".to_string()); + socketeer.send(message.clone()).await.unwrap(); + let received_message = socketeer.next_message().await.unwrap(); + assert_eq!(message, received_message); + socketeer = socketeer.reconnect().await.unwrap(); + let message = EchoControlMessage::Message("Hello".to_string()); + socketeer.send(message.clone()).await.unwrap(); + let received_message = socketeer.next_message().await.unwrap(); + assert_eq!(message, received_message); + socketeer.close_connection().await.unwrap(); +} + +#[tokio::test] +async fn test_closed_socket() { + let server_address = get_mock_address(echo_server).await; + let mut socketeer: Socketeer = Socketeer::connect(&format!("ws://{server_address}")) + .await + .unwrap(); + let close_request = EchoControlMessage::Close; + socketeer.send(close_request.clone()).await.unwrap(); + let response = socketeer.next_message().await; + assert!(matches!(response.unwrap_err(), Error::WebsocketClosed)); + let send_result = socketeer.send(close_request).await; + assert!(send_result.is_err()); + let error = send_result.unwrap_err(); + println!("Actual Error: {error:#?}"); + assert!(matches!(error, Error::WebsocketClosed)); +} + +#[tokio::test] +async fn test_close_request() { + let server_address = get_mock_address(echo_server).await; + let socketeer: Socketeer = Socketeer::connect(&format!("ws://{server_address}")) + .await + .unwrap(); + socketeer.close_connection().await.unwrap(); +} + +#[tokio::test] +async fn test_connect_with_default_options() { + let server_address = get_mock_address(echo_server).await; + let mut socketeer: Socketeer = + Socketeer::connect_with(&format!("ws://{server_address}"), ConnectOptions::default()) + .await + .unwrap(); + let message = EchoControlMessage::Message("Hello".to_string()); + socketeer.send(message.clone()).await.unwrap(); + let received_message = socketeer.next_message().await.unwrap(); + assert_eq!(message, received_message); +} + +#[tokio::test] +async fn test_raw_codec_message_roundtrip() { + // Typed send/next_message round-trip when the codec is RawCodec — the + // codec is identity, so frames pass through unchanged. + let server_address = get_mock_address(echo_server).await; + let mut socketeer: Socketeer = Socketeer::connect(&format!("ws://{server_address}")) + .await + .unwrap(); + let raw_text = r#"{"Message":"raw hello"}"#; + socketeer + .send(Message::Text(raw_text.into())) + .await + .unwrap(); + let received = socketeer.next_message().await.unwrap(); + assert_eq!(received, Message::Text(raw_text.into())); +} + +#[tokio::test] +async fn test_disabled_keepalive() { + let server_address = get_mock_address(echo_server).await; + let options = ConnectOptions { + keepalive_interval: None, + ..ConnectOptions::default() + }; + let mut socketeer: Socketeer = + Socketeer::connect_with(&format!("ws://{server_address}"), options) + .await + .unwrap(); + let message = EchoControlMessage::Message("Hello".to_string()); + socketeer.send(message.clone()).await.unwrap(); + let received_message = socketeer.next_message().await.unwrap(); + assert_eq!(message, received_message); +} + +#[tokio::test] +async fn test_handler_on_connected() { + use serde::{Deserialize, Serialize}; + use std::sync::Arc; + use tokio::sync::Mutex; + + #[derive(Clone, Debug, Serialize, Deserialize, PartialEq)] + struct AuthResponse { + status: String, + } + + struct TestAuthHandler { + connected_count: Arc>, + } + + impl ConnectionHandler for TestAuthHandler { + async fn on_connected(&mut self, ctx: &mut HandshakeContext<'_, C>) -> Result<(), Error> { + ctx.send_text(r#"{"action":"auth","token":"test-token"}"#) + .await?; + let text = ctx.recv_text().await?; + let response: AuthResponse = serde_json::from_str(&text).unwrap(); + assert_eq!(response.status, "authenticated"); + let mut count = self.connected_count.lock().await; + *count += 1; + Ok(()) + } + } + + let connected_count = Arc::new(Mutex::new(0u32)); + let handler = TestAuthHandler { + connected_count: connected_count.clone(), + }; + + let server_address = get_mock_address(auth_echo_server).await; + let mut socketeer: Socketeer = Socketeer::connect_with_codec( + &format!("ws://{server_address}"), + ConnectOptions::default(), + JsonCodec::new(), + handler, + ) + .await + .unwrap(); + + assert_eq!(*connected_count.lock().await, 1); + + let message = EchoControlMessage::Message("after auth".to_string()); + socketeer.send(message.clone()).await.unwrap(); + let received = socketeer.next_message().await.unwrap(); + assert_eq!(message, received); +} + +#[tokio::test] +async fn test_handler_reconnect() { + use std::sync::Arc; + use tokio::sync::Mutex; + + struct ReconnectHandler { + connected_count: Arc>, + disconnected_count: Arc>, + } + + impl ConnectionHandler for ReconnectHandler { + async fn on_connected(&mut self, ctx: &mut HandshakeContext<'_, C>) -> Result<(), Error> { + ctx.send_text(r#"{"action":"auth","token":"test-token"}"#) + .await?; + let _response = ctx.recv_text().await?; + let mut count = self.connected_count.lock().await; + *count += 1; + Ok(()) + } + + async fn on_disconnected(&mut self) { + let mut count = self.disconnected_count.lock().await; + *count += 1; + } + } + + let connected_count = Arc::new(Mutex::new(0u32)); + let disconnected_count = Arc::new(Mutex::new(0u32)); + let handler = ReconnectHandler { + connected_count: connected_count.clone(), + disconnected_count: disconnected_count.clone(), + }; + + let server_address = get_mock_address(auth_echo_server).await; + let mut socketeer = Socketeer::::connect_with_codec( + &format!("ws://{server_address}"), + ConnectOptions::default(), + JsonCodec::new(), + handler, + ) + .await + .unwrap(); + + assert_eq!(*connected_count.lock().await, 1); + assert_eq!(*disconnected_count.lock().await, 0); + + // Send a message to verify connection works + let message = EchoControlMessage::Message("before reconnect".to_string()); + socketeer.send(message.clone()).await.unwrap(); + let received = socketeer.next_message().await.unwrap(); + assert_eq!(message, received); + + // Reconnect — handler should fire again + socketeer = socketeer.reconnect().await.unwrap(); + + assert_eq!(*connected_count.lock().await, 2); + assert_eq!(*disconnected_count.lock().await, 1); + + // Verify connection still works after reconnect + let message = EchoControlMessage::Message("after reconnect".to_string()); + socketeer.send(message.clone()).await.unwrap(); + let received = socketeer.next_message().await.unwrap(); + assert_eq!(message, received); + + socketeer.close_connection().await.unwrap(); +} + +#[cfg(feature = "msgpack")] +#[tokio::test] +async fn test_msgpack_send_receive() { + type EchoMsgPack = MsgPackCodec; + + let server_address = get_mock_address(msgpack_echo_server).await; + let mut socketeer: Socketeer = + Socketeer::connect(&format!("ws://{server_address}")) + .await + .unwrap(); + let message = EchoControlMessage::Message("msgpack hello".to_string()); + socketeer.send(message.clone()).await.unwrap(); + let received = socketeer.next_message().await.unwrap(); + assert_eq!(message, received); + socketeer.close_connection().await.unwrap(); +} + +#[tokio::test] +async fn test_handler_uses_codec_driven_send_recv() { + // Exercises HandshakeContext::send / recv (the codec-driven path). + // Other handler tests only cover the raw send_text / recv_text helpers. + struct TypedHandshakeHandler; + + impl ConnectionHandler for TypedHandshakeHandler { + async fn on_connected( + &mut self, + ctx: &mut HandshakeContext<'_, EchoJson>, + ) -> Result<(), Error> { + ctx.send(&EchoControlMessage::Message("handshake".into())) + .await?; + let echoed = ctx.recv().await?; + assert_eq!(echoed, EchoControlMessage::Message("handshake".into())); + Ok(()) + } + } + + let server_address = get_mock_address(echo_server).await; + let mut socketeer: Socketeer = Socketeer::connect_with_codec( + &format!("ws://{server_address}"), + ConnectOptions::default(), + JsonCodec::new(), + TypedHandshakeHandler, + ) + .await + .unwrap(); + + // Confirm normal traffic still flows after the typed handshake. + let message = EchoControlMessage::Message("after handshake".into()); + socketeer.send(message.clone()).await.unwrap(); + assert_eq!(socketeer.next_message().await.unwrap(), message); + socketeer.close_connection().await.unwrap(); +} + +#[tokio::test] +async fn test_handshake_recv_close_with_raw_codec() { + // Regression: with RawCodec, recv_raw returns Ok(Message::Close(_)) and + // RawCodec::decode is the identity, so a peer-initiated close used to + // surface as Ok(Close) instead of Err(WebsocketClosed). recv must + // intercept Close before delegating to the codec. + struct CloseExpecting; + + impl ConnectionHandler for CloseExpecting { + async fn on_connected( + &mut self, + ctx: &mut HandshakeContext<'_, RawCodec>, + ) -> Result<(), Error> { + // Ask the echo server to close (JSON unit-variant for EchoControlMessage::Close). + ctx.send(&Message::Text(r#""Close""#.into())).await?; + let err = ctx.recv().await.unwrap_err(); + assert!(matches!(err, Error::WebsocketClosed)); + Ok(()) + } + } + + let server_address = get_mock_address(echo_server).await; + let _socketeer: Socketeer = Socketeer::connect_with_codec( + &format!("ws://{server_address}"), + ConnectOptions::default(), + RawCodec::new(), + CloseExpecting, + ) + .await + .unwrap(); +} + +#[tokio::test] +async fn test_extra_headers_used() { + // Cover ConnectOptions::build_request's loop body that copies + // `extra_headers` onto the upgrade request. + let server_address = get_mock_address(echo_server).await; + let mut headers = tokio_tungstenite::tungstenite::http::HeaderMap::new(); + headers.insert("X-Test-Header", "socketeer".parse().unwrap()); + let options = ConnectOptions { + extra_headers: headers, + ..ConnectOptions::default() + }; + let mut socketeer: Socketeer = + Socketeer::connect_with(&format!("ws://{server_address}"), options) + .await + .unwrap(); + let message = EchoControlMessage::Message("hi".into()); + socketeer.send(message.clone()).await.unwrap(); + assert_eq!(socketeer.next_message().await.unwrap(), message); + socketeer.close_connection().await.unwrap(); +} + +#[tokio::test] +async fn test_auth_handler_bad_token() { + // Covers auth_echo_server's bad-token branch (sends {"status":"error"} + // and shuts down). The handler observes the error response, returns + // Ok, then a subsequent send fails because the server has closed. + struct BadTokenHandler; + + impl ConnectionHandler for BadTokenHandler { + async fn on_connected(&mut self, ctx: &mut HandshakeContext<'_, C>) -> Result<(), Error> { + ctx.send_text(r#"{"action":"auth","token":"WRONG"}"#) + .await?; + let resp = ctx.recv_text().await?; + assert!(resp.contains("error")); + Ok(()) + } + } + + let server_address = get_mock_address(auth_echo_server).await; + let _socketeer: Socketeer = Socketeer::connect_with_codec( + &format!("ws://{server_address}"), + ConnectOptions::default(), + JsonCodec::new(), + BadTokenHandler, + ) + .await + .unwrap(); +} + +#[cfg(feature = "msgpack")] +#[tokio::test] +async fn test_msgpack_send_ping() { + // Covers the SendPing arm of msgpack_echo_server. + type EchoMsgPack = MsgPackCodec; + + let server_address = get_mock_address(msgpack_echo_server).await; + let mut socketeer: Socketeer = + Socketeer::connect(&format!("ws://{server_address}")) + .await + .unwrap(); + socketeer.send(EchoControlMessage::SendPing).await.unwrap(); + // Server replies with a Ping; Socketeer auto-Pongs. Round-trip a real + // message to confirm the connection is still alive. + let message = EchoControlMessage::Message("after ping".into()); + socketeer.send(message.clone()).await.unwrap(); + assert_eq!(socketeer.next_message().await.unwrap(), message); + socketeer.close_connection().await.unwrap(); +} + +#[cfg(feature = "msgpack")] +#[tokio::test] +async fn test_msgpack_close_request() { + // Covers the Close arm of msgpack_echo_server. + type EchoMsgPack = MsgPackCodec; + + let server_address = get_mock_address(msgpack_echo_server).await; + let mut socketeer: Socketeer = + Socketeer::connect(&format!("ws://{server_address}")) + .await + .unwrap(); + socketeer.send(EchoControlMessage::Close).await.unwrap(); + let result = socketeer.next_message().await; + assert!(matches!(result.unwrap_err(), Error::WebsocketClosed)); +} + +#[tokio::test] +async fn test_socketeer_debug_format() { + let server_address = get_mock_address(echo_server).await; + let socketeer: Socketeer = Socketeer::connect(&format!("ws://{server_address}")) + .await + .unwrap(); + let formatted = format!("{socketeer:?}"); + assert!(formatted.starts_with("Socketeer")); + assert!(formatted.contains("url")); +} + +#[tokio::test] +async fn test_send_raw_next_raw_message() { + // Cover the raw send/receive escape hatches on a typed (non-RawCodec) + // connection: send_raw bypasses encoding, next_raw_message bypasses + // decoding, so we can speak frames the codec wouldn't otherwise + // produce or accept. + let server_address = get_mock_address(echo_server).await; + let mut socketeer: Socketeer = Socketeer::connect(&format!("ws://{server_address}")) + .await + .unwrap(); + let raw_text = r#"{"Message":"raw recv"}"#; + socketeer + .send_raw(Message::Text(raw_text.into())) + .await + .unwrap(); + let frame = socketeer.next_raw_message().await.unwrap(); + assert_eq!(frame, Message::Text(raw_text.into())); + socketeer.close_connection().await.unwrap(); +} + +#[cfg(feature = "msgpack")] +#[tokio::test] +async fn test_handshake_send_binary_recv_raw() { + // Cover HandshakeContext::send_binary by sending a pre-encoded + // msgpack frame from on_connected and reading the binary echo back + // via recv_raw. + struct BinaryHandshake; + + type EchoMsgPack = MsgPackCodec; + + impl ConnectionHandler for BinaryHandshake { + async fn on_connected( + &mut self, + ctx: &mut HandshakeContext<'_, EchoMsgPack>, + ) -> Result<(), Error> { + let payload = + rmp_serde::to_vec_named(&EchoControlMessage::Message("binary".into())).unwrap(); + ctx.send_binary(payload).await?; + let echo = ctx.recv_raw().await?; + assert!(matches!(echo, Message::Binary(_))); + Ok(()) + } + } + + let server_address = get_mock_address(msgpack_echo_server).await; + let socketeer: Socketeer = Socketeer::connect_with_codec( + &format!("ws://{server_address}"), + ConnectOptions::default(), + MsgPackCodec::new(), + BinaryHandshake, + ) + .await + .unwrap(); + socketeer.close_connection().await.unwrap(); +} + +#[cfg(feature = "msgpack")] +#[tokio::test] +async fn test_handshake_recv_text_rejects_binary() { + // Cover the non-Text branch of HandshakeContext::recv_text by pointing + // it at a server that only speaks binary frames. + struct ExpectsTextOnBinary; + + type EchoMsgPack = MsgPackCodec; + + impl ConnectionHandler for ExpectsTextOnBinary { + async fn on_connected( + &mut self, + ctx: &mut HandshakeContext<'_, EchoMsgPack>, + ) -> Result<(), Error> { + let payload = + rmp_serde::to_vec_named(&EchoControlMessage::Message("hi".into())).unwrap(); + ctx.send_binary(payload).await?; + // recv_text must reject the echoed Binary frame. + let err = ctx.recv_text().await.unwrap_err(); + assert!(matches!(err, Error::UnexpectedMessageType(_))); + Ok(()) + } + } + + let server_address = get_mock_address(msgpack_echo_server).await; + let socketeer: Socketeer = Socketeer::connect_with_codec( + &format!("ws://{server_address}"), + ConnectOptions::default(), + MsgPackCodec::new(), + ExpectsTextOnBinary, + ) + .await + .unwrap(); + socketeer.close_connection().await.unwrap(); +} + +#[tokio::test] +async fn test_binary_custom_keepalive() { + // The widening of custom_keepalive_message from Option to + // Option is otherwise unexercised. echo_server silently + // ignores Binary frames, so the receive queue stays clean and we can + // verify the connection survives a binary keepalive cycle. + let server_address = get_mock_address(echo_server).await; + let options = ConnectOptions { + keepalive_interval: Some(Duration::from_millis(100)), + custom_keepalive_message: Some(Message::Binary(Bytes::from_static(b"keepalive"))), + ..ConnectOptions::default() + }; + let mut socketeer: Socketeer = + Socketeer::connect_with(&format!("ws://{server_address}"), options) + .await + .unwrap(); + + // Wait long enough for at least a couple of keepalive ticks to fire. + sleep(Duration::from_millis(350)).await; + + let message = EchoControlMessage::Message("post-keepalive".into()); + socketeer.send(message.clone()).await.unwrap(); + assert_eq!(socketeer.next_message().await.unwrap(), message); + socketeer.close_connection().await.unwrap(); +}