1use futures::{SinkExt, StreamExt};
40use std::io::{Error as IoError, ErrorKind};
41use tokio::time::{Duration, timeout};
42use tokio_tungstenite::{
43 MaybeTlsStream, WebSocketStream, connect_async,
44 tungstenite::{Error as WsError, Message},
45};
46
47pub struct WebSocketTestClient {
52 stream: WebSocketStream<MaybeTlsStream<tokio::net::TcpStream>>,
54 url: String,
56}
57
58impl WebSocketTestClient {
59 pub async fn connect(url: &str) -> Result<Self, WsError> {
73 let (stream, _response) = connect_async(url).await?;
74 Ok(Self {
75 stream,
76 url: url.to_string(),
77 })
78 }
79
80 pub async fn connect_with_token(url: &str, token: &str) -> Result<Self, WsError> {
99 use tokio_tungstenite::tungstenite::http::Request;
100
101 let request = Request::builder()
102 .uri(url)
103 .header("Authorization", format!("Bearer {}", token))
104 .body(())
105 .expect("Failed to build WebSocket request");
106
107 let (stream, _response) = connect_async(request).await?;
108 Ok(Self {
109 stream,
110 url: url.to_string(),
111 })
112 }
113
114 pub async fn connect_with_query_token(url: &str, token: &str) -> Result<Self, WsError> {
134 let url_with_token = format!("{}?token={}", url, urlencoding::encode(token));
135 Self::connect(&url_with_token).await
136 }
137
138 pub async fn connect_with_cookie(
158 url: &str,
159 cookie_name: &str,
160 cookie_value: &str,
161 ) -> Result<Self, WsError> {
162 use tokio_tungstenite::tungstenite::http::Request;
163
164 let request = Request::builder()
165 .uri(url)
166 .header("Cookie", format!("{}={}", cookie_name, cookie_value))
167 .body(())
168 .expect("Failed to build WebSocket request");
169
170 let (stream, _response) = connect_async(request).await?;
171 Ok(Self {
172 stream,
173 url: url.to_string(),
174 })
175 }
176
177 pub async fn send_text(&mut self, text: &str) -> Result<(), WsError> {
192 self.stream.send(Message::text(text)).await
193 }
194
195 pub async fn send_binary(&mut self, data: &[u8]) -> Result<(), WsError> {
197 self.stream.send(Message::binary(data.to_vec())).await
198 }
199
200 pub async fn send_ping(&mut self, payload: &[u8]) -> Result<(), WsError> {
202 self.stream
203 .send(Message::Ping(payload.to_vec().into()))
204 .await
205 }
206
207 pub async fn send_pong(&mut self, payload: &[u8]) -> Result<(), WsError> {
209 self.stream
210 .send(Message::Pong(payload.to_vec().into()))
211 .await
212 }
213
214 pub async fn receive(&mut self) -> Option<Result<Message, WsError>> {
218 self.stream.next().await
219 }
220
221 pub async fn receive_text(&mut self) -> Result<String, WsError> {
237 self.receive_text_with_timeout(Duration::from_secs(5)).await
238 }
239
240 pub async fn receive_text_with_timeout(
242 &mut self,
243 duration: Duration,
244 ) -> Result<String, WsError> {
245 match timeout(duration, self.stream.next()).await {
246 Ok(Some(Ok(Message::Text(text)))) => Ok(text.to_string()),
247 Ok(Some(Ok(msg))) => Err(WsError::Io(IoError::new(
248 ErrorKind::InvalidData,
249 format!("Expected text message, got {:?}", msg),
250 ))),
251 Ok(Some(Err(e))) => Err(e),
252 Ok(None) => Err(WsError::ConnectionClosed),
253 Err(_) => Err(WsError::Io(IoError::new(
254 ErrorKind::TimedOut,
255 "Receive timeout",
256 ))),
257 }
258 }
259
260 pub async fn receive_binary(&mut self) -> Result<Vec<u8>, WsError> {
262 self.receive_binary_with_timeout(Duration::from_secs(5))
263 .await
264 }
265
266 pub async fn receive_binary_with_timeout(
268 &mut self,
269 duration: Duration,
270 ) -> Result<Vec<u8>, WsError> {
271 match timeout(duration, self.stream.next()).await {
272 Ok(Some(Ok(Message::Binary(data)))) => Ok(data.to_vec()),
273 Ok(Some(Ok(msg))) => Err(WsError::Io(IoError::new(
274 ErrorKind::InvalidData,
275 format!("Expected binary message, got {:?}", msg),
276 ))),
277 Ok(Some(Err(e))) => Err(e),
278 Ok(None) => Err(WsError::ConnectionClosed),
279 Err(_) => Err(WsError::Io(IoError::new(
280 ErrorKind::TimedOut,
281 "Receive timeout",
282 ))),
283 }
284 }
285
286 pub async fn close(mut self) -> Result<(), WsError> {
288 self.stream.close(None).await
289 }
290
291 pub fn url(&self) -> &str {
293 &self.url
294 }
295}
296
297pub mod assertions {
299 use tokio_tungstenite::tungstenite::Message;
300
301 pub fn assert_message_text(msg: &Message, expected: &str) {
312 match msg {
313 Message::Text(text) => assert_eq!(text.as_str(), expected),
314 _ => panic!("Expected text message, got {:?}", msg),
315 }
316 }
317
318 pub fn assert_message_contains(msg: &Message, substring: &str) {
320 match msg {
321 Message::Text(text) => assert!(
322 text.contains(substring),
323 "Message '{}' does not contain '{}'",
324 text,
325 substring
326 ),
327 _ => panic!("Expected text message, got {:?}", msg),
328 }
329 }
330
331 pub fn assert_message_binary(msg: &Message, expected: &[u8]) {
333 match msg {
334 Message::Binary(data) => assert_eq!(data.as_ref(), expected),
335 _ => panic!("Expected binary message, got {:?}", msg),
336 }
337 }
338
339 pub fn assert_message_ping(msg: &Message) {
341 match msg {
342 Message::Ping(_) => {}
343 _ => panic!("Expected ping message, got {:?}", msg),
344 }
345 }
346
347 pub fn assert_message_pong(msg: &Message) {
349 match msg {
350 Message::Pong(_) => {}
351 _ => panic!("Expected pong message, got {:?}", msg),
352 }
353 }
354}
355
356#[cfg(test)]
357mod tests {
358 use super::*;
359 use std::sync::{Arc, Mutex};
360
361 use rstest::rstest;
362 use tokio::net::TcpListener;
363 use tokio::sync::Notify;
364 use tokio::task::JoinHandle;
365 use tokio_tungstenite::accept_async;
366
367 struct WebSocketServerGuard {
368 handle: Option<JoinHandle<()>>,
369 }
370
371 impl WebSocketServerGuard {
372 async fn join(mut self) {
373 if let Some(handle) = self.handle.take() {
374 handle.await.unwrap();
375 }
376 }
377 }
378
379 impl Drop for WebSocketServerGuard {
380 fn drop(&mut self) {
381 if let Some(handle) = self.handle.take() {
382 handle.abort();
383 }
384 }
385 }
386
387 async fn start_echo_server(
388 initial_message: Option<Message>,
389 close_messages: Arc<Mutex<Vec<Message>>>,
390 ) -> (String, WebSocketServerGuard) {
391 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
392 let address = listener.local_addr().unwrap();
393 let handle = tokio::spawn(async move {
394 let (stream, _) = listener.accept().await.unwrap();
395 let mut socket = accept_async(stream).await.unwrap();
396 if let Some(message) = initial_message {
397 socket.send(message).await.unwrap();
398 }
399 while let Some(result) = socket.next().await {
400 let message = result.unwrap();
401 match message {
402 Message::Close(frame) => {
403 close_messages.lock().unwrap().push(Message::Close(frame));
404 break;
405 }
406 Message::Text(text) => socket.send(Message::Text(text)).await.unwrap(),
407 Message::Binary(bytes) => socket.send(Message::Binary(bytes)).await.unwrap(),
408 Message::Ping(bytes) => socket.send(Message::Pong(bytes)).await.unwrap(),
409 Message::Pong(bytes) => socket.send(Message::Ping(bytes)).await.unwrap(),
410 Message::Frame(_) => unreachable!("raw frames are not yielded by tungstenite"),
411 }
412 }
413 });
414
415 (
416 format!("ws://{address}/"),
417 WebSocketServerGuard {
418 handle: Some(handle),
419 },
420 )
421 }
422
423 async fn start_timeout_server() -> (String, WebSocketServerGuard, Arc<Notify>) {
424 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
425 let address = listener.local_addr().unwrap();
426 let release = Arc::new(Notify::new());
427 let server_release = Arc::clone(&release);
428 let handle = tokio::spawn(async move {
429 let (stream, _) = listener.accept().await.unwrap();
430 let _socket = accept_async(stream).await.unwrap();
431 server_release.notified().await;
432 });
433
434 (
435 format!("ws://{address}"),
436 WebSocketServerGuard {
437 handle: Some(handle),
438 },
439 release,
440 )
441 }
442
443 #[test]
444 fn test_url_with_query_token() {
445 let url = "ws://localhost:8080/ws";
446 let token = "my-token";
447 let expected = "ws://localhost:8080/ws?token=my-token";
448
449 let url_with_token = format!("{}?token={}", url, urlencoding::encode(token));
450 assert_eq!(url_with_token, expected);
451 }
452
453 #[test]
454 fn test_url_with_query_token_special_chars() {
455 let url = "ws://localhost:8080/ws";
456 let token = "token with spaces&special=chars";
457 let url_with_token = format!("{}?token={}", url, urlencoding::encode(token));
458 assert_eq!(
459 url_with_token,
460 "ws://localhost:8080/ws?token=token%20with%20spaces%26special%3Dchars"
461 );
462 }
463
464 #[test]
465 fn test_message_assertions() {
466 use assertions::*;
467
468 let text_msg = Message::text("Hello");
469 assert_message_text(&text_msg, "Hello");
470 assert_message_contains(&text_msg, "ell");
471
472 let binary_msg = Message::Binary(vec![1, 2, 3].into());
473 assert_message_binary(&binary_msg, &[1, 2, 3]);
474
475 let ping_msg = Message::Ping(vec![].into());
476 assert_message_ping(&ping_msg);
477
478 let pong_msg = Message::Pong(vec![].into());
479 assert_message_pong(&pong_msg);
480 }
481
482 #[rstest]
483 #[tokio::test]
484 async fn websocket_client_roundtrips_frames_and_reports_timeout_and_type_errors() {
485 let close_messages = Arc::new(Mutex::new(Vec::new()));
487 let (url, echo_server) = start_echo_server(None, Arc::clone(&close_messages)).await;
488 let (wrong_url, wrong_server) = start_echo_server(
489 Some(Message::binary(&b"not text"[..])),
490 Arc::new(Mutex::new(Vec::new())),
491 )
492 .await;
493 let (timeout_url, timeout_server, timeout_release) = start_timeout_server().await;
494 let mut client = WebSocketTestClient::connect(&url).await.unwrap();
495 let client_url = client.url().to_string();
496 let mut wrong_client = WebSocketTestClient::connect(&wrong_url).await.unwrap();
497 let mut timeout_client = WebSocketTestClient::connect(&timeout_url).await.unwrap();
498
499 client.send_text("hello websocket").await.unwrap();
501 let text = client.receive_text().await.unwrap();
502 client.send_binary(&[1, 2, 3]).await.unwrap();
503 let binary = client.receive_binary().await.unwrap();
504 client.send_ping(b"ping-data").await.unwrap();
505 let pong = client.receive().await.unwrap().unwrap();
506 client.send_pong(b"pong-data").await.unwrap();
507 let ping = client.receive().await.unwrap().unwrap();
508 let wrong_frame = wrong_client
509 .receive_text_with_timeout(Duration::from_millis(50))
510 .await
511 .unwrap_err();
512 let timeout_error = timeout_client
513 .receive_text_with_timeout(Duration::from_millis(10))
514 .await
515 .unwrap_err();
516 timeout_release.notify_one();
517 wrong_client.close().await.unwrap();
518 client.close().await.unwrap();
519
520 assert_eq!(client_url, url);
522 assert_eq!(text, "hello websocket");
523 assert_eq!(binary, vec![1, 2, 3]);
524 assert_eq!(pong, Message::Pong(b"ping-data".to_vec().into()));
525 assert_eq!(ping, Message::Ping(b"pong-data".to_vec().into()));
526 match wrong_frame {
527 WsError::Io(error) => {
528 assert_eq!(error.kind(), ErrorKind::InvalidData);
529 assert_eq!(
530 error.to_string(),
531 "Expected text message, got Binary(b\"not text\")"
532 );
533 }
534 other => panic!("expected invalid data error, got {other:?}"),
535 }
536 match timeout_error {
537 WsError::Io(error) => {
538 assert_eq!(error.kind(), ErrorKind::TimedOut);
539 assert_eq!(error.to_string(), "Receive timeout");
540 }
541 other => panic!("expected timeout error, got {other:?}"),
542 }
543 echo_server.join().await;
544 wrong_server.join().await;
545 timeout_server.join().await;
546 assert_eq!(*close_messages.lock().unwrap(), vec![Message::Close(None)]);
547 }
548
549 #[tokio::test]
550 async fn websocket_client_query_auth_encodes_token_in_connected_url() {
551 let close_messages = Arc::new(Mutex::new(Vec::new()));
553 let (url, server) = start_echo_server(None, Arc::clone(&close_messages)).await;
554
555 let query_client = WebSocketTestClient::connect_with_query_token(&url, "space & equals=")
557 .await
558 .unwrap();
559 let query_url = query_client.url().to_string();
560 query_client.close().await.unwrap();
561 server.join().await;
562
563 assert_eq!(query_url, format!("{url}?token=space%20%26%20equals%3D"));
565 assert_eq!(*close_messages.lock().unwrap(), vec![Message::Close(None)]);
566 }
567}