1use std::{convert::Infallible, time::Duration};
2
3use rama_core::{
4 Layer, Service,
5 bytes::Bytes,
6 error::{BoxError, ErrorExt},
7 extensions::{self, Extensions, ExtensionsRef},
8 futures::{
9 Sink, SinkExt as _, Stream, StreamExt as _,
10 channel::{mpsc, oneshot},
11 },
12 io::{BridgeIo, Io},
13 service::MirrorService,
14 telemetry::tracing,
15};
16
17use crate::{
18 AsyncWebSocket, ProtocolError, Utf8Bytes, WebSocketIo,
19 handshake::matcher::RelayWebSocketConfig,
20 protocol::{CloseFrame, Role, frame::coding::CloseCode},
21};
22
23const DEFAULT_CLOSE_HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(5);
24
25#[derive(Debug)]
35pub struct WebSocketBridge<Ingress, Egress> {
36 pub ingress: Ingress,
38 pub egress: Egress,
40}
41
42#[derive(Debug, Clone)]
50pub struct WebSocketRelayIoService<S> {
51 inner: S,
52}
53
54#[derive(Debug, Clone, Copy, Default)]
57pub struct WebSocketRelayIoLayer;
58
59impl WebSocketRelayIoLayer {
60 #[must_use]
62 pub const fn new() -> Self {
63 Self
64 }
65}
66
67impl<S> Layer<S> for WebSocketRelayIoLayer {
68 type Service = WebSocketRelayIoService<S>;
69
70 fn layer(&self, inner: S) -> Self::Service {
71 WebSocketRelayIoService::new(inner)
72 }
73}
74
75impl<S> WebSocketRelayIoService<S> {
76 #[must_use]
78 pub const fn new(inner: S) -> Self {
79 Self { inner }
80 }
81
82 #[must_use]
84 pub const fn inner(&self) -> &S {
85 &self.inner
86 }
87
88 #[must_use]
90 pub fn into_inner(self) -> S {
91 self.inner
92 }
93}
94
95impl<S, Ingress, Egress> Service<BridgeIo<Ingress, Egress>> for WebSocketRelayIoService<S>
96where
97 S: Service<WebSocketBridge<AsyncWebSocket<Ingress>, AsyncWebSocket<Egress>>>,
98 Ingress: Io + Unpin + ExtensionsRef,
99 Egress: Io + Unpin + ExtensionsRef,
100{
101 type Output = S::Output;
102 type Error = S::Error;
103
104 async fn serve(&self, bridge: BridgeIo<Ingress, Egress>) -> Result<Self::Output, Self::Error> {
105 self.inner
106 .serve(upgrade_websocket_bridge(bridge).await)
107 .await
108 }
109}
110
111#[derive(Debug, Clone)]
112pub struct WebSocketRelayService<S = MirrorService> {
142 middleware: S,
143 close_handshake_timeout: Duration,
144}
145
146impl<S> WebSocketRelayService<S> {
147 #[inline(always)]
148 #[must_use]
149 pub fn new(middleware: S) -> Self {
151 Self {
152 middleware,
153 close_handshake_timeout: DEFAULT_CLOSE_HANDSHAKE_TIMEOUT,
154 }
155 }
156
157 rama_utils::macros::generate_set_and_with! {
158 pub fn close_handshake_timeout(mut self, timeout: Duration) -> Self {
164 self.close_handshake_timeout = timeout;
165 self
166 }
167 }
168}
169
170#[derive(Debug, Clone)]
171pub struct WebSocketRelayEventService<S = MirrorService> {
189 middleware: S,
190 close_handshake_timeout: Duration,
191}
192
193impl<S> WebSocketRelayEventService<S> {
194 #[inline(always)]
195 #[must_use]
196 pub fn new(middleware: S) -> Self {
198 Self {
199 middleware,
200 close_handshake_timeout: DEFAULT_CLOSE_HANDSHAKE_TIMEOUT,
201 }
202 }
203
204 rama_utils::macros::generate_set_and_with! {
205 pub fn close_handshake_timeout(mut self, timeout: Duration) -> Self {
211 self.close_handshake_timeout = timeout;
212 self
213 }
214 }
215}
216
217#[derive(Debug, Clone)]
218pub struct WebSocketRelayInput {
221 pub direction: WebSocketRelayDirection,
222 pub message: WebSocketRelayMessage,
223 pub extensions: Extensions,
224}
225
226impl ExtensionsRef for WebSocketRelayInput {
227 #[inline(always)]
228 fn extensions(&self) -> &Extensions {
229 &self.extensions
230 }
231}
232
233#[derive(Debug, Clone)]
234pub struct WebSocketRelayOutput {
237 pub messages: Vec<WebSocketRelayMessage>,
241 pub extensions: Extensions,
245}
246
247impl From<WebSocketRelayInput> for WebSocketRelayOutput {
248 fn from(value: WebSocketRelayInput) -> Self {
249 let WebSocketRelayInput {
250 direction: _,
251 message,
252 extensions,
253 } = value;
254
255 Self {
256 messages: vec![message],
257 extensions,
258 }
259 }
260}
261
262impl ExtensionsRef for WebSocketRelayOutput {
263 #[inline(always)]
264 fn extensions(&self) -> &Extensions {
265 &self.extensions
266 }
267}
268
269#[derive(Debug, Clone)]
270pub struct WebSocketRelayEventInput {
272 pub direction: WebSocketRelayDirection,
273 pub event: WebSocketRelayEvent,
274 pub extensions: Extensions,
275}
276
277impl ExtensionsRef for WebSocketRelayEventInput {
278 #[inline(always)]
279 fn extensions(&self) -> &Extensions {
280 &self.extensions
281 }
282}
283
284#[derive(Debug, Clone)]
285pub struct WebSocketRelayEventOutput {
287 pub messages: Vec<WebSocketRelayMessage>,
291 pub close: Option<WebSocketRelayClose>,
297 pub extensions: Extensions,
301}
302
303impl From<WebSocketRelayEventInput> for WebSocketRelayEventOutput {
304 fn from(value: WebSocketRelayEventInput) -> Self {
305 let WebSocketRelayEventInput {
306 direction: _,
307 event,
308 extensions,
309 } = value;
310
311 let messages = match event {
312 WebSocketRelayEvent::Data(message) => vec![message],
313 WebSocketRelayEvent::Ping(_)
314 | WebSocketRelayEvent::Pong(_)
315 | WebSocketRelayEvent::Close(_) => Vec::new(),
316 };
317
318 Self {
319 messages,
320 close: None,
321 extensions,
322 }
323 }
324}
325
326impl ExtensionsRef for WebSocketRelayEventOutput {
327 #[inline(always)]
328 fn extensions(&self) -> &Extensions {
329 &self.extensions
330 }
331}
332
333#[derive(Debug, Clone, Eq, PartialEq)]
334pub enum WebSocketRelayEvent {
339 Data(WebSocketRelayMessage),
341 Ping(Bytes),
343 Pong(Bytes),
345 Close(Option<CloseFrame>),
347}
348
349#[derive(Debug, Clone, Eq, PartialEq)]
350pub enum WebSocketRelayClose {
355 WithoutFrame,
357 WithFrame(CloseFrame),
363}
364
365impl From<Option<CloseFrame>> for WebSocketRelayClose {
366 fn from(value: Option<CloseFrame>) -> Self {
367 match value {
368 Some(frame) => Self::WithFrame(frame),
369 None => Self::WithoutFrame,
370 }
371 }
372}
373
374impl WebSocketRelayClose {
375 fn into_frame(self) -> Option<CloseFrame> {
376 match self {
377 Self::WithoutFrame => None,
378 Self::WithFrame(frame) => Some(frame),
379 }
380 }
381}
382
383#[derive(Debug, Clone, Eq, PartialEq)]
384pub enum WebSocketRelayMessage {
387 Text(Utf8Bytes),
389 Binary(Bytes),
391}
392
393impl From<WebSocketRelayMessage> for crate::protocol::Message {
394 fn from(value: WebSocketRelayMessage) -> Self {
395 match value {
396 WebSocketRelayMessage::Text(utf8_bytes) => Self::Text(utf8_bytes),
397 WebSocketRelayMessage::Binary(bytes) => Self::Binary(bytes),
398 }
399 }
400}
401
402#[derive(Debug, Clone, Copy, PartialEq, Eq)]
403pub enum WebSocketRelayDirection {
406 Ingress,
407 Egress,
408}
409
410impl<S, Ingress, Egress> Service<BridgeIo<Ingress, Egress>> for WebSocketRelayService<S>
411where
412 S: Service<WebSocketRelayInput, Output: Into<WebSocketRelayOutput>, Error: Into<BoxError>>,
413 Ingress: Io + Unpin + extensions::ExtensionsRef,
414 Egress: Io + Unpin + extensions::ExtensionsRef,
415{
416 type Output = ();
417 type Error = Infallible;
418
419 async fn serve(&self, bridge: BridgeIo<Ingress, Egress>) -> Result<Self::Output, Self::Error> {
420 let WebSocketBridge {
421 ingress: ingress_socket,
422 egress: egress_socket,
423 } = upgrade_websocket_bridge(bridge).await;
424 relay_websocket_bridge(
425 MessageRelayHandler {
426 middleware: &self.middleware,
427 },
428 ingress_socket,
429 egress_socket,
430 self.close_handshake_timeout,
431 )
432 .await;
433 Ok(())
434 }
435}
436
437impl<S, Ingress, Egress> Service<WebSocketBridge<Ingress, Egress>> for WebSocketRelayService<S>
438where
439 S: Service<WebSocketRelayInput, Output: Into<WebSocketRelayOutput>, Error: Into<BoxError>>,
440 Ingress: WebSocketIo,
441 Egress: WebSocketIo,
442{
443 type Output = ();
444 type Error = Infallible;
445
446 async fn serve(
447 &self,
448 WebSocketBridge {
449 ingress: ingress_socket,
450 egress: egress_socket,
451 }: WebSocketBridge<Ingress, Egress>,
452 ) -> Result<Self::Output, Self::Error> {
453 relay_websocket_bridge(
454 MessageRelayHandler {
455 middleware: &self.middleware,
456 },
457 ingress_socket,
458 egress_socket,
459 self.close_handshake_timeout,
460 )
461 .await;
462 Ok(())
463 }
464}
465
466impl<S, Ingress, Egress> Service<BridgeIo<Ingress, Egress>> for WebSocketRelayEventService<S>
467where
468 S: Service<
469 WebSocketRelayEventInput,
470 Output: Into<WebSocketRelayEventOutput>,
471 Error: Into<BoxError>,
472 >,
473 Ingress: Io + Unpin + extensions::ExtensionsRef,
474 Egress: Io + Unpin + extensions::ExtensionsRef,
475{
476 type Output = ();
477 type Error = Infallible;
478
479 async fn serve(&self, bridge: BridgeIo<Ingress, Egress>) -> Result<Self::Output, Self::Error> {
480 let WebSocketBridge {
481 ingress: ingress_socket,
482 egress: egress_socket,
483 } = upgrade_websocket_bridge(bridge).await;
484 relay_websocket_bridge(
485 EventRelayHandler {
486 middleware: &self.middleware,
487 },
488 ingress_socket,
489 egress_socket,
490 self.close_handshake_timeout,
491 )
492 .await;
493 Ok(())
494 }
495}
496
497impl<S, Ingress, Egress> Service<WebSocketBridge<Ingress, Egress>> for WebSocketRelayEventService<S>
498where
499 S: Service<
500 WebSocketRelayEventInput,
501 Output: Into<WebSocketRelayEventOutput>,
502 Error: Into<BoxError>,
503 >,
504 Ingress: WebSocketIo,
505 Egress: WebSocketIo,
506{
507 type Output = ();
508 type Error = Infallible;
509
510 async fn serve(
511 &self,
512 WebSocketBridge {
513 ingress: ingress_socket,
514 egress: egress_socket,
515 }: WebSocketBridge<Ingress, Egress>,
516 ) -> Result<Self::Output, Self::Error> {
517 relay_websocket_bridge(
518 EventRelayHandler {
519 middleware: &self.middleware,
520 },
521 ingress_socket,
522 egress_socket,
523 self.close_handshake_timeout,
524 )
525 .await;
526 Ok(())
527 }
528}
529
530struct RelayHandlerOutput {
531 messages: Vec<WebSocketRelayMessage>,
532 close: Option<WebSocketRelayClose>,
533 extensions: Extensions,
534}
535
536trait RelayHandler {
537 fn serve(
538 &self,
539 direction: WebSocketRelayDirection,
540 event: WebSocketRelayEvent,
541 extensions: Extensions,
542 ) -> impl Future<Output = Result<RelayHandlerOutput, BoxError>> + Send + '_;
543}
544
545struct MessageRelayHandler<'a, S> {
546 middleware: &'a S,
547}
548
549impl<S> RelayHandler for MessageRelayHandler<'_, S>
550where
551 S: Service<WebSocketRelayInput, Output: Into<WebSocketRelayOutput>, Error: Into<BoxError>>,
552{
553 async fn serve(
554 &self,
555 direction: WebSocketRelayDirection,
556 event: WebSocketRelayEvent,
557 extensions: Extensions,
558 ) -> Result<RelayHandlerOutput, BoxError> {
559 let WebSocketRelayEvent::Data(message) = event else {
560 return Ok(RelayHandlerOutput {
561 messages: Vec::new(),
562 close: None,
563 extensions,
564 });
565 };
566
567 let WebSocketRelayOutput {
568 messages,
569 extensions,
570 } = self
571 .middleware
572 .serve(WebSocketRelayInput {
573 direction,
574 message,
575 extensions,
576 })
577 .await
578 .map(Into::into)
579 .map_err(Into::into)?;
580
581 Ok(RelayHandlerOutput {
582 messages,
583 close: None,
584 extensions,
585 })
586 }
587}
588
589struct EventRelayHandler<'a, S> {
590 middleware: &'a S,
591}
592
593impl<S> RelayHandler for EventRelayHandler<'_, S>
594where
595 S: Service<
596 WebSocketRelayEventInput,
597 Output: Into<WebSocketRelayEventOutput>,
598 Error: Into<BoxError>,
599 >,
600{
601 async fn serve(
602 &self,
603 direction: WebSocketRelayDirection,
604 event: WebSocketRelayEvent,
605 extensions: Extensions,
606 ) -> Result<RelayHandlerOutput, BoxError> {
607 let WebSocketRelayEventOutput {
608 messages,
609 close,
610 extensions,
611 } = self
612 .middleware
613 .serve(WebSocketRelayEventInput {
614 direction,
615 event,
616 extensions,
617 })
618 .await
619 .map(Into::into)
620 .map_err(Into::into)?;
621
622 Ok(RelayHandlerOutput {
623 messages,
624 close,
625 extensions,
626 })
627 }
628}
629
630async fn upgrade_websocket_bridge<Ingress, Egress>(
631 BridgeIo(ingress_stream, egress_stream): BridgeIo<Ingress, Egress>,
632) -> WebSocketBridge<AsyncWebSocket<Ingress>, AsyncWebSocket<Egress>>
633where
634 Ingress: Io + Unpin + ExtensionsRef,
635 Egress: Io + Unpin + ExtensionsRef,
636{
637 let maybe_ws_config = egress_stream
638 .extensions()
639 .get_ref()
640 .map(|RelayWebSocketConfig(cfg)| *cfg);
641
642 let ingress_socket =
643 AsyncWebSocket::from_raw_socket(ingress_stream, Role::Server, maybe_ws_config).await;
644 let egress_socket =
645 AsyncWebSocket::from_raw_socket(egress_stream, Role::Client, maybe_ws_config).await;
646 WebSocketBridge {
647 ingress: ingress_socket,
648 egress: egress_socket,
649 }
650}
651
652async fn relay_websocket_bridge<H, Ingress, Egress>(
653 handler: H,
654 ingress_socket: Ingress,
655 egress_socket: Egress,
656 close_handshake_timeout: Duration,
657) where
658 H: RelayHandler,
659 Ingress: WebSocketIo,
660 Egress: WebSocketIo,
661{
662 let mut ingress_relay_extensions = ingress_socket.extensions().fork();
666 let mut egress_relay_extensions = egress_socket.extensions().fork();
667
668 let (ingress_writer, ingress_reader) = ingress_socket.split();
669 let (egress_writer, egress_reader) = egress_socket.split();
670
671 let (ingress_writer_tx, ingress_writer_rx) = mpsc::unbounded();
672 let (egress_writer_tx, egress_writer_rx) = mpsc::unbounded();
673 let (ingress_close_tx, ingress_close_rx) = mpsc::unbounded();
674 let (egress_close_tx, egress_close_rx) = mpsc::unbounded();
675 let close_controls = CloseControls {
676 ingress: ingress_close_tx,
677 egress: egress_close_tx,
678 };
679 let (signal_tx, signal_rx) = mpsc::unbounded();
680
681 let ingress_direction = relay_direction(
682 &handler,
683 WebSocketRelayDirection::Ingress,
684 ingress_reader,
685 DirectionChannels {
686 source_writer: ingress_writer_tx.clone(),
687 destination_writer: egress_writer_tx.clone(),
688 close_control: ingress_close_rx,
689 close_controls: close_controls.clone(),
690 signals: signal_tx.clone(),
691 },
692 &mut ingress_relay_extensions,
693 );
694 let egress_direction = relay_direction(
695 &handler,
696 WebSocketRelayDirection::Egress,
697 egress_reader,
698 DirectionChannels {
699 source_writer: egress_writer_tx,
700 destination_writer: ingress_writer_tx,
701 close_control: egress_close_rx,
702 close_controls,
703 signals: signal_tx,
704 },
705 &mut egress_relay_extensions,
706 );
707 let drivers = async {
708 tokio::join!(
709 writer_loop("ingress", ingress_writer, ingress_writer_rx),
710 writer_loop("egress", egress_writer, egress_writer_rx),
711 ingress_direction,
712 egress_direction,
713 );
714 };
715
716 tokio::select! {
717 () = supervise_relay(signal_rx, close_handshake_timeout) => {}
718 () = drivers => {}
719 }
720}
721
722#[derive(Debug)]
723enum WriterCommand {
724 Send {
725 message: crate::Message,
726 response: oneshot::Sender<Result<(), ProtocolError>>,
727 },
728 Flush {
729 response: oneshot::Sender<Result<(), ProtocolError>>,
730 },
731}
732
733async fn writer_loop<Socket>(
734 socket_name: &'static str,
735 mut socket: Socket,
736 mut commands: mpsc::UnboundedReceiver<WriterCommand>,
737) where
738 Socket: Sink<crate::Message, Error = ProtocolError> + Unpin,
739{
740 while let Some(command) = commands.next().await {
741 let (result, response) = match command {
742 WriterCommand::Send { message, response } => (socket.send(message).await, response),
743 WriterCommand::Flush { response } => (socket.flush().await, response),
744 };
745 if response.send(result).is_err() {
746 tracing::trace!("{socket_name} WS writer response receiver was dropped");
747 }
748 }
749}
750
751fn queue_message(
752 writer: &mpsc::UnboundedSender<WriterCommand>,
753 message: crate::Message,
754) -> Option<oneshot::Receiver<Result<(), ProtocolError>>> {
755 let (response, receiver) = oneshot::channel();
756 writer
757 .unbounded_send(WriterCommand::Send { message, response })
758 .ok()
759 .map(|()| receiver)
760}
761
762fn queue_flush(
763 writer: &mpsc::UnboundedSender<WriterCommand>,
764) -> Option<oneshot::Receiver<Result<(), ProtocolError>>> {
765 let (response, receiver) = oneshot::channel();
766 writer
767 .unbounded_send(WriterCommand::Flush { response })
768 .ok()
769 .map(|()| receiver)
770}
771
772#[derive(Clone)]
773struct CloseControls {
774 ingress: mpsc::UnboundedSender<()>,
775 egress: mpsc::UnboundedSender<()>,
776}
777
778impl CloseControls {
779 fn start_closing(&self) {
780 _ = self.ingress.unbounded_send(());
781 _ = self.egress.unbounded_send(());
782 }
783}
784
785#[derive(Debug, Clone, Copy)]
786enum RelaySignal {
787 ClosingStarted,
788 SideFinished(WebSocketRelayDirection),
789 Terminate,
790}
791
792fn signal(signals: &mpsc::UnboundedSender<RelaySignal>, signal: RelaySignal) {
793 _ = signals.unbounded_send(signal);
794}
795
796async fn supervise_relay(
797 mut signals: mpsc::UnboundedReceiver<RelaySignal>,
798 close_handshake_timeout: Duration,
799) {
800 let mut ingress_finished = false;
801 let mut egress_finished = false;
802
803 loop {
804 match signals.next().await {
805 Some(RelaySignal::ClosingStarted) => break,
806 Some(RelaySignal::SideFinished(WebSocketRelayDirection::Ingress)) => {
807 ingress_finished = true;
808 }
809 Some(RelaySignal::SideFinished(WebSocketRelayDirection::Egress)) => {
810 egress_finished = true;
811 }
812 Some(RelaySignal::Terminate) | None => return,
813 }
814 }
815
816 let finish_close = async {
817 while !ingress_finished || !egress_finished {
818 match signals.next().await {
819 Some(RelaySignal::SideFinished(WebSocketRelayDirection::Ingress)) => {
820 ingress_finished = true;
821 }
822 Some(RelaySignal::SideFinished(WebSocketRelayDirection::Egress)) => {
823 egress_finished = true;
824 }
825 Some(RelaySignal::ClosingStarted) => {}
826 Some(RelaySignal::Terminate) | None => return,
827 }
828 }
829 };
830
831 if tokio::time::timeout(close_handshake_timeout, finish_close)
832 .await
833 .is_err()
834 {
835 tracing::debug!(
836 ?close_handshake_timeout,
837 "WS close handshake timed out; drop MITM relay"
838 );
839 }
840}
841
842enum WriterWait {
843 Complete(Result<(), ProtocolError>),
844 StartClosing,
845}
846
847struct DirectionChannels {
848 source_writer: mpsc::UnboundedSender<WriterCommand>,
849 destination_writer: mpsc::UnboundedSender<WriterCommand>,
850 close_control: mpsc::UnboundedReceiver<()>,
851 close_controls: CloseControls,
852 signals: mpsc::UnboundedSender<RelaySignal>,
853}
854
855async fn await_writer_or_close(
856 response: oneshot::Receiver<Result<(), ProtocolError>>,
857 close_control: &mut mpsc::UnboundedReceiver<()>,
858) -> WriterWait {
859 tokio::select! {
860 biased;
861 _ = close_control.next() => WriterWait::StartClosing,
862 result = response => match result {
863 Ok(result) => WriterWait::Complete(result),
864 Err(_) => WriterWait::Complete(Err(ProtocolError::Io(std::io::Error::new(
865 std::io::ErrorKind::ConnectionAborted,
866 "WS writer task ended",
867 )))),
868 }
869 }
870}
871
872async fn relay_direction<H, Source>(
873 handler: &H,
874 direction: WebSocketRelayDirection,
875 mut source: Source,
876 channels: DirectionChannels,
877 relay_extensions: &mut Extensions,
878) where
879 H: RelayHandler,
880 Source: Stream<Item = Result<crate::Message, ProtocolError>> + Unpin,
881{
882 let DirectionChannels {
883 source_writer,
884 destination_writer,
885 mut close_control,
886 close_controls,
887 signals,
888 } = channels;
889 let (source_name, destination_name) = match direction {
890 WebSocketRelayDirection::Ingress => ("ingress", "egress"),
891 WebSocketRelayDirection::Egress => ("egress", "ingress"),
892 };
893
894 loop {
895 let source_result = tokio::select! {
896 biased;
897 _ = close_control.next() => {
898 return drain_close(
899 direction,
900 source_name,
901 &mut source,
902 &source_writer,
903 &signals,
904 ).await;
905 }
906 result = source.next() => result,
907 };
908
909 let message = match source_result {
910 Some(Ok(message)) => message,
911 Some(Err(error)) => {
912 tracing::debug!(
913 "{source_name} WS socket ended with protocol error ({error})... drop MITM relay"
914 );
915 signal(&signals, RelaySignal::Terminate);
916 return;
917 }
918 None => {
919 tracing::debug!("{source_name} WS socket disconnected... drop MITM relay");
920 signal(&signals, RelaySignal::Terminate);
921 return;
922 }
923 };
924
925 let (event, flush_automatic_response, opposite_heartbeat) = match message {
926 crate::Message::Text(text) => (
927 WebSocketRelayEvent::Data(WebSocketRelayMessage::Text(text)),
928 false,
929 None,
930 ),
931 crate::Message::Binary(bytes) => (
932 WebSocketRelayEvent::Data(WebSocketRelayMessage::Binary(bytes)),
933 false,
934 None,
935 ),
936 crate::Message::Ping(bytes) => {
937 (WebSocketRelayEvent::Ping(bytes.clone()), true, Some(bytes))
938 }
939 crate::Message::Pong(bytes) => (WebSocketRelayEvent::Pong(bytes), false, None),
940 crate::Message::Close(frame) => {
941 let event = WebSocketRelayEvent::Close(frame.clone());
942 let flush = queue_flush(&source_writer);
943 if queue_message(&destination_writer, crate::Message::Close(frame)).is_none() {
944 tracing::debug!("failed to queue close for {destination_name} WS socket");
945 }
946 close_controls.start_closing();
947 signal(&signals, RelaySignal::ClosingStarted);
948
949 let observe = async {
950 match handler
951 .serve(direction, event, std::mem::take(relay_extensions))
952 .await
953 {
954 Ok(output) => {
955 tracing::trace!(
956 discarded_message_count = output.messages.len(),
957 discarded_close_request = output.close.is_some(),
958 "ignore WS relay middleware output returned while observing {source_name} close"
959 );
960 }
961 Err(error) => {
962 tracing::debug!(
963 "WS relay middleware failed while observing {source_name} close: ({})...",
964 error.into_box_error()
965 );
966 }
967 }
968 };
969 let complete = finish_close_side(direction, source_name, flush, &signals);
970 tokio::join!(observe, complete);
971 return;
972 }
973 crate::Message::Frame(_) => {
974 tracing::debug!(
975 "unexpected raw frame returned while reading {source_name} WS socket; drop it"
976 );
977 continue;
978 }
979 };
980
981 if flush_automatic_response {
982 let Some(response) = queue_flush(&source_writer) else {
983 tracing::debug!(
984 "failed to queue automatic WS control response for {source_name} socket"
985 );
986 signal(&signals, RelaySignal::Terminate);
987 return;
988 };
989 match await_writer_or_close(response, &mut close_control).await {
990 WriterWait::Complete(Ok(())) => {}
991 WriterWait::Complete(Err(error)) => {
992 tracing::debug!(
993 "failed to flush automatic WS control response to {source_name} socket: {error}; drop MITM relay"
994 );
995 signal(&signals, RelaySignal::Terminate);
996 return;
997 }
998 WriterWait::StartClosing => {
999 return drain_close(
1000 direction,
1001 source_name,
1002 &mut source,
1003 &source_writer,
1004 &signals,
1005 )
1006 .await;
1007 }
1008 }
1009 }
1010
1011 if let Some(payload) = opposite_heartbeat {
1012 let Some(response) = queue_message(&destination_writer, crate::Message::Pong(payload))
1016 else {
1017 tracing::debug!("failed to queue WS heartbeat for {destination_name} socket");
1018 signal(&signals, RelaySignal::Terminate);
1019 return;
1020 };
1021 match await_writer_or_close(response, &mut close_control).await {
1022 WriterWait::Complete(Ok(())) => {}
1023 WriterWait::Complete(Err(error)) => {
1024 tracing::debug!(
1025 "failed to send WS heartbeat to {destination_name}: {error}; drop MITM relay"
1026 );
1027 signal(&signals, RelaySignal::Terminate);
1028 return;
1029 }
1030 WriterWait::StartClosing => {
1031 return drain_close(
1032 direction,
1033 source_name,
1034 &mut source,
1035 &source_writer,
1036 &signals,
1037 )
1038 .await;
1039 }
1040 }
1041 }
1042
1043 let extensions = std::mem::take(relay_extensions);
1044 let handler_result = tokio::select! {
1045 biased;
1046 _ = close_control.next() => {
1047 return drain_close(
1048 direction,
1049 source_name,
1050 &mut source,
1051 &source_writer,
1052 &signals,
1053 ).await;
1054 }
1055 result = handler.serve(direction, event, extensions) => result,
1056 };
1057
1058 let RelayHandlerOutput {
1059 messages,
1060 close,
1061 extensions,
1062 } = match handler_result {
1063 Ok(output) => output,
1064 Err(error) => {
1065 tracing::debug!(
1066 "WS relay middleware failed on {source_name} event: ({})... close both connections",
1067 error.into_box_error()
1068 );
1069 start_coordinated_close(
1070 &source_writer,
1071 &destination_writer,
1072 &close_controls,
1073 &signals,
1074 Some(internal_error_close("relay middleware error")),
1075 );
1076 return drain_close(
1077 direction,
1078 source_name,
1079 &mut source,
1080 &source_writer,
1081 &signals,
1082 )
1083 .await;
1084 }
1085 };
1086 *relay_extensions = extensions;
1087
1088 let requested_close = close.map(WebSocketRelayClose::into_frame);
1089 if requested_close
1090 .as_ref()
1091 .is_some_and(|frame| !valid_close_frame(frame.as_ref()))
1092 {
1093 tracing::debug!(
1094 "WS relay middleware returned an invalid close frame on {source_name} event; close both connections with 1011"
1095 );
1096 start_coordinated_close(
1097 &source_writer,
1098 &destination_writer,
1099 &close_controls,
1100 &signals,
1101 Some(internal_error_close("invalid relay close frame")),
1102 );
1103 return drain_close(
1104 direction,
1105 source_name,
1106 &mut source,
1107 &source_writer,
1108 &signals,
1109 )
1110 .await;
1111 }
1112
1113 for (message_index, message) in messages.into_iter().enumerate() {
1114 tracing::trace!(
1115 "relay {source_name} WS data message #{message_index} to {destination_name}"
1116 );
1117 let Some(response) = queue_message(&destination_writer, message.into()) else {
1118 tracing::debug!(
1119 "{destination_name} WS writer ended @ message#{message_index}; drop MITM relay"
1120 );
1121 signal(&signals, RelaySignal::Terminate);
1122 return;
1123 };
1124 match await_writer_or_close(response, &mut close_control).await {
1125 WriterWait::Complete(Ok(())) => {}
1126 WriterWait::Complete(Err(error)) => {
1127 tracing::debug!(
1128 "failed to relay {source_name} message to {destination_name}: {error} @ message#{message_index}; drop MITM relay"
1129 );
1130 signal(&signals, RelaySignal::Terminate);
1131 return;
1132 }
1133 WriterWait::StartClosing => {
1134 return drain_close(
1135 direction,
1136 source_name,
1137 &mut source,
1138 &source_writer,
1139 &signals,
1140 )
1141 .await;
1142 }
1143 }
1144 }
1145
1146 if let Some(frame) = requested_close {
1147 start_coordinated_close(
1148 &source_writer,
1149 &destination_writer,
1150 &close_controls,
1151 &signals,
1152 frame,
1153 );
1154 return drain_close(
1155 direction,
1156 source_name,
1157 &mut source,
1158 &source_writer,
1159 &signals,
1160 )
1161 .await;
1162 }
1163 }
1164}
1165
1166fn start_coordinated_close(
1167 source_writer: &mpsc::UnboundedSender<WriterCommand>,
1168 destination_writer: &mpsc::UnboundedSender<WriterCommand>,
1169 close_controls: &CloseControls,
1170 signals: &mpsc::UnboundedSender<RelaySignal>,
1171 frame: Option<CloseFrame>,
1172) {
1173 _ = queue_message(source_writer, crate::Message::Close(frame.clone()));
1174 _ = queue_message(destination_writer, crate::Message::Close(frame));
1175 close_controls.start_closing();
1176 signal(signals, RelaySignal::ClosingStarted);
1177}
1178
1179async fn drain_close<Source>(
1180 direction: WebSocketRelayDirection,
1181 source_name: &'static str,
1182 source: &mut Source,
1183 source_writer: &mpsc::UnboundedSender<WriterCommand>,
1184 signals: &mpsc::UnboundedSender<RelaySignal>,
1185) where
1186 Source: Stream<Item = Result<crate::Message, ProtocolError>> + Unpin,
1187{
1188 loop {
1189 match source.next().await {
1190 Some(Ok(crate::Message::Close(_))) => {
1191 finish_close_side(direction, source_name, queue_flush(source_writer), signals)
1192 .await;
1193 return;
1194 }
1195 Some(Ok(_)) => {}
1196 Some(Err(error)) => {
1197 tracing::debug!(
1198 "{source_name} WS socket ended while closing with protocol error: {error}"
1199 );
1200 signal(signals, RelaySignal::SideFinished(direction));
1201 return;
1202 }
1203 None => {
1204 tracing::debug!("{source_name} WS socket disconnected while closing");
1205 signal(signals, RelaySignal::SideFinished(direction));
1206 return;
1207 }
1208 }
1209 }
1210}
1211
1212async fn finish_close_side(
1213 direction: WebSocketRelayDirection,
1214 source_name: &'static str,
1215 response: Option<oneshot::Receiver<Result<(), ProtocolError>>>,
1216 signals: &mpsc::UnboundedSender<RelaySignal>,
1217) {
1218 match response {
1219 Some(response) => match response.await {
1220 Ok(Ok(())) => {}
1221 Ok(Err(error)) => {
1222 tracing::debug!(
1223 "failed to flush automatic close response to {source_name} WS socket: {error}"
1224 );
1225 }
1226 Err(_) => {
1227 tracing::debug!("{source_name} WS writer ended while flushing close response");
1228 }
1229 },
1230 None => {
1231 tracing::debug!("failed to queue close response flush for {source_name} WS socket");
1232 }
1233 }
1234 signal(signals, RelaySignal::SideFinished(direction));
1235}
1236
1237fn valid_close_frame(frame: Option<&CloseFrame>) -> bool {
1238 frame.is_none_or(|frame| frame.code.is_allowed() && frame.reason.len() <= 123)
1239}
1240
1241fn internal_error_close(reason: &'static str) -> CloseFrame {
1242 CloseFrame {
1243 code: CloseCode::Error,
1244 reason: reason.into(),
1245 }
1246}
1247
1248#[cfg(test)]
1249mod tests {
1250 use parking_lot::Mutex;
1256 use std::{future::pending, sync::Arc, time::Duration};
1257
1258 use rama_core::{
1259 Layer, Service,
1260 bytes::Bytes,
1261 error::{BoxError, BoxErrorExt as _},
1262 extensions::{Extension, Extensions, ExtensionsRef},
1263 futures::{SinkExt as _, channel::oneshot},
1264 io::{BridgeIo, Io},
1265 service::MirrorService,
1266 };
1267 use rama_net::test_utils::client::MockSocket;
1268 use rama_utils::octets::kib;
1269 use tokio::{io::duplex, time::timeout};
1270
1271 use crate::{
1272 AsyncWebSocket, Message,
1273 handshake::mitm::{
1274 WebSocketBridge, WebSocketRelayClose, WebSocketRelayDirection, WebSocketRelayEvent,
1275 WebSocketRelayEventInput, WebSocketRelayEventOutput, WebSocketRelayEventService,
1276 WebSocketRelayInput, WebSocketRelayIoLayer, WebSocketRelayIoService,
1277 WebSocketRelayMessage, WebSocketRelayOutput, WebSocketRelayService, valid_close_frame,
1278 },
1279 protocol::{CloseFrame, Role, frame::coding::CloseCode},
1280 };
1281
1282 #[derive(Debug, Clone, Extension)]
1283 struct IngressMarker;
1284
1285 #[derive(Debug, Clone, Extension)]
1286 struct EgressMarker;
1287
1288 #[derive(Debug, Clone, Extension)]
1289 struct LeakProbeIngress;
1290
1291 #[derive(Debug, Clone, Extension)]
1292 struct LeakProbeEgress;
1293
1294 #[test]
1295 fn message_bridge_and_raw_adapter_preserve_their_inputs() {
1296 let ingress = Extensions::new();
1297 ingress.insert(IngressMarker);
1298 let bridge = WebSocketBridge {
1299 ingress,
1300 egress: Extensions::new(),
1301 };
1302 assert!(
1303 bridge
1304 .ingress
1305 .extensions()
1306 .get_ref::<IngressMarker>()
1307 .is_some()
1308 );
1309 assert!(
1310 bridge
1311 .egress
1312 .extensions()
1313 .get_ref::<IngressMarker>()
1314 .is_none()
1315 );
1316
1317 let adapter = WebSocketRelayIoService::new(42_u8);
1318 assert_eq!(adapter.inner(), &42);
1319 assert_eq!(adapter.into_inner(), 42);
1320
1321 let adapter = WebSocketRelayIoLayer::new().into_layer(42_u8);
1322 assert_eq!(adapter.into_inner(), 42);
1323 }
1324
1325 #[derive(Debug, Clone)]
1326 struct Observation {
1327 direction: WebSocketRelayDirection,
1328 saw_ingress_marker: bool,
1329 saw_egress_marker: bool,
1330 saw_leak_ingress: bool,
1331 saw_leak_egress: bool,
1332 }
1333
1334 #[derive(Clone)]
1335 struct RecordingMiddleware {
1336 log: Arc<Mutex<Vec<Observation>>>,
1337 }
1338
1339 impl Service<WebSocketRelayInput> for RecordingMiddleware {
1340 type Output = WebSocketRelayOutput;
1341 type Error = BoxError;
1342
1343 async fn serve(&self, input: WebSocketRelayInput) -> Result<Self::Output, Self::Error> {
1344 let WebSocketRelayInput {
1345 direction,
1346 message,
1347 extensions,
1348 } = input;
1349 let obs = Observation {
1350 direction,
1351 saw_ingress_marker: extensions.get_ref::<IngressMarker>().is_some(),
1354 saw_egress_marker: extensions.get_ref::<EgressMarker>().is_some(),
1355 saw_leak_ingress: extensions.get_ref::<LeakProbeIngress>().is_some(),
1358 saw_leak_egress: extensions.get_ref::<LeakProbeEgress>().is_some(),
1359 };
1360 self.log.lock().push(obs);
1361 match direction {
1362 WebSocketRelayDirection::Ingress => {
1363 extensions.insert(LeakProbeIngress);
1364 }
1365 WebSocketRelayDirection::Egress => {
1366 extensions.insert(LeakProbeEgress);
1367 }
1368 }
1369 Ok(WebSocketRelayOutput {
1370 messages: vec![message],
1371 extensions,
1372 })
1373 }
1374 }
1375
1376 #[tokio::test]
1377 async fn relay_per_direction_fork_isolation() {
1378 let (relay_ingress_dup, peer_ingress_dup) = duplex(kib(16));
1382 let (relay_egress_dup, peer_egress_dup) = duplex(kib(16));
1383
1384 let relay_ingress = MockSocket::new(relay_ingress_dup);
1385 let relay_egress = MockSocket::new(relay_egress_dup);
1386 relay_ingress.extensions().insert(IngressMarker);
1387 relay_egress.extensions().insert(EgressMarker);
1388
1389 let ingress_live_ext = relay_ingress.extensions().clone();
1396 let egress_live_ext = relay_egress.extensions().clone();
1397
1398 let log = Arc::new(Mutex::new(Vec::<Observation>::new()));
1399 let middleware = RecordingMiddleware { log: log.clone() };
1400 let svc = WebSocketRelayService::new(middleware);
1401
1402 let relay =
1403 tokio::spawn(async move { svc.serve(BridgeIo(relay_ingress, relay_egress)).await });
1404
1405 let peer_ingress = MockSocket::new(peer_ingress_dup);
1408 let mut peer_ingress_ws =
1409 AsyncWebSocket::from_raw_socket(peer_ingress, Role::Client, None).await;
1410 let peer_egress = MockSocket::new(peer_egress_dup);
1411 let mut peer_egress_ws =
1412 AsyncWebSocket::from_raw_socket(peer_egress, Role::Server, None).await;
1413
1414 peer_ingress_ws
1416 .send_message(Message::text("ping"))
1417 .await
1418 .expect("peer ingress send");
1419 match expect_message(&mut peer_egress_ws, "peer egress recv").await {
1420 Message::Text(t) => assert_eq!(t.as_str(), "ping"),
1421 other => panic!("unexpected message on egress peer: {other:?}"),
1422 }
1423
1424 peer_egress_ws
1426 .send_message(Message::text("pong"))
1427 .await
1428 .expect("peer egress send");
1429 match expect_message(&mut peer_ingress_ws, "peer ingress recv").await {
1430 Message::Text(t) => assert_eq!(t.as_str(), "pong"),
1431 other => panic!("unexpected message on ingress peer: {other:?}"),
1432 }
1433
1434 drop(peer_ingress_ws);
1437 drop(peer_egress_ws);
1438 _ = relay.await.expect("relay task join");
1439
1440 let log = log.lock();
1441 assert_eq!(log.len(), 2, "exactly one middleware call per direction");
1442
1443 let ingress = log
1444 .iter()
1445 .find(|o| o.direction == WebSocketRelayDirection::Ingress)
1446 .expect("ingress observation");
1447 let egress = log
1448 .iter()
1449 .find(|o| o.direction == WebSocketRelayDirection::Egress)
1450 .expect("egress observation");
1451
1452 assert!(
1454 ingress.saw_ingress_marker,
1455 "ingress fork sees IngressMarker"
1456 );
1457 assert!(egress.saw_egress_marker, "egress fork sees EgressMarker");
1458 assert!(
1459 !ingress.saw_egress_marker,
1460 "ingress fork must NOT see EgressMarker (forks are independent)"
1461 );
1462 assert!(
1463 !egress.saw_ingress_marker,
1464 "egress fork must NOT see IngressMarker (forks are independent)"
1465 );
1466
1467 assert!(
1472 !ingress.saw_leak_egress,
1473 "ingress fork must NOT see LeakProbeEgress (cross-direction leak)"
1474 );
1475 assert!(
1476 !egress.saw_leak_ingress,
1477 "egress fork must NOT see LeakProbeIngress (cross-direction leak)"
1478 );
1479
1480 assert!(
1484 !ingress_live_ext.self_contains::<LeakProbeIngress>(),
1485 "LeakProbeIngress must NOT leak onto the live ingress socket"
1486 );
1487 assert!(
1488 !egress_live_ext.self_contains::<LeakProbeEgress>(),
1489 "LeakProbeEgress must NOT leak onto the live egress socket"
1490 );
1491 }
1492
1493 async fn expect_message<Stream>(
1494 socket: &mut AsyncWebSocket<Stream>,
1495 description: &str,
1496 ) -> Message
1497 where
1498 Stream: Io + Unpin,
1499 {
1500 timeout(Duration::from_secs(1), socket.recv_message())
1501 .await
1502 .unwrap_or_else(|_| panic!("timed out waiting to {description}"))
1503 .unwrap_or_else(|error| panic!("failed to {description}: {error}"))
1504 }
1505
1506 async fn assert_no_message<Stream>(socket: &mut AsyncWebSocket<Stream>, peer_name: &str)
1507 where
1508 Stream: Io + Unpin,
1509 {
1510 assert!(
1511 timeout(Duration::from_millis(50), socket.recv_message())
1512 .await
1513 .is_err(),
1514 "{peer_name} unexpectedly received a message"
1515 );
1516 }
1517
1518 fn test_close_frame(reason: &'static str) -> CloseFrame {
1519 CloseFrame {
1520 code: CloseCode::Away,
1521 reason: reason.into(),
1522 }
1523 }
1524
1525 #[tokio::test]
1526 async fn regular_relay_keeps_both_legs_active_when_either_peer_pings() {
1527 let (relay_ingress_dup, peer_ingress_dup) = duplex(kib(16));
1528 let (relay_egress_dup, peer_egress_dup) = duplex(kib(16));
1529
1530 let log = Arc::new(Mutex::new(Vec::<Observation>::new()));
1531 let service = WebSocketRelayService::new(RecordingMiddleware { log: log.clone() });
1532 let relay = tokio::spawn(async move {
1533 service
1534 .serve(BridgeIo(
1535 MockSocket::new(relay_ingress_dup),
1536 MockSocket::new(relay_egress_dup),
1537 ))
1538 .await
1539 });
1540
1541 let mut peer_ingress_ws =
1542 AsyncWebSocket::from_raw_socket(MockSocket::new(peer_ingress_dup), Role::Client, None)
1543 .await;
1544 let mut peer_egress_ws =
1545 AsyncWebSocket::from_raw_socket(MockSocket::new(peer_egress_dup), Role::Server, None)
1546 .await;
1547
1548 let ingress_ping = Bytes::from_static(b"ingress-ping");
1549 peer_ingress_ws
1550 .send_message(Message::Ping(ingress_ping.clone()))
1551 .await
1552 .expect("send ingress ping");
1553 assert_eq!(
1554 expect_message(&mut peer_ingress_ws, "receive automatic ingress pong").await,
1555 Message::Pong(ingress_ping)
1556 );
1557 assert_no_message(
1558 &mut peer_ingress_ws,
1559 "ingress peer after its automatic pong",
1560 )
1561 .await;
1562 assert_eq!(
1563 expect_message(&mut peer_egress_ws, "receive egress heartbeat").await,
1564 Message::Pong(Bytes::from_static(b"ingress-ping"))
1565 );
1566
1567 let unsolicited_pong = Bytes::from_static(b"unsolicited-pong");
1568 peer_ingress_ws
1569 .send_message(Message::Pong(unsolicited_pong))
1570 .await
1571 .expect("send unsolicited ingress pong");
1572 assert_no_message(&mut peer_egress_ws, "egress peer after ingress pong").await;
1573
1574 let egress_ping = Bytes::from_static(b"egress-ping");
1575 peer_egress_ws
1576 .send_message(Message::Ping(egress_ping.clone()))
1577 .await
1578 .expect("send egress ping");
1579 assert_eq!(
1580 expect_message(&mut peer_egress_ws, "receive automatic egress pong").await,
1581 Message::Pong(egress_ping)
1582 );
1583 assert_no_message(&mut peer_egress_ws, "egress peer after its automatic pong").await;
1584 assert_eq!(
1585 expect_message(&mut peer_ingress_ws, "receive ingress heartbeat").await,
1586 Message::Pong(Bytes::from_static(b"egress-ping"))
1587 );
1588
1589 drop(peer_ingress_ws);
1590 drop(peer_egress_ws);
1591 _ = relay.await.expect("relay task join");
1592 assert!(
1593 log.lock().is_empty(),
1594 "regular middleware must not observe ping or pong"
1595 );
1596 }
1597
1598 #[tokio::test]
1599 async fn regular_relay_coordinates_both_close_handshakes() {
1600 let (relay_ingress_dup, peer_ingress_dup) = duplex(kib(16));
1601 let (relay_egress_dup, peer_egress_dup) = duplex(kib(16));
1602
1603 let service = WebSocketRelayService::new(MirrorService::new());
1604 let mut relay = tokio::spawn(async move {
1605 service
1606 .serve(BridgeIo(
1607 MockSocket::new(relay_ingress_dup),
1608 MockSocket::new(relay_egress_dup),
1609 ))
1610 .await
1611 });
1612
1613 let mut peer_ingress_ws =
1614 AsyncWebSocket::from_raw_socket(MockSocket::new(peer_ingress_dup), Role::Client, None)
1615 .await;
1616 let mut peer_egress_ws =
1617 AsyncWebSocket::from_raw_socket(MockSocket::new(peer_egress_dup), Role::Server, None)
1618 .await;
1619
1620 let close_frame = test_close_frame("regular shutdown");
1621 peer_ingress_ws
1622 .send_message(Message::Close(Some(close_frame.clone())))
1623 .await
1624 .expect("send ingress close");
1625
1626 assert_eq!(
1627 expect_message(&mut peer_ingress_ws, "receive ingress close reply").await,
1628 Message::Close(Some(close_frame.clone()))
1629 );
1630 assert_eq!(
1631 expect_message(&mut peer_egress_ws, "receive propagated egress close").await,
1632 Message::Close(Some(close_frame))
1633 );
1634
1635 assert!(
1636 timeout(Duration::from_millis(50), &mut relay)
1637 .await
1638 .is_err(),
1639 "relay must await the egress peer's close reply"
1640 );
1641 peer_egress_ws
1642 .flush()
1643 .await
1644 .expect("flush egress close reply");
1645 timeout(Duration::from_secs(1), relay)
1646 .await
1647 .expect("relay close timeout")
1648 .expect("relay task join")
1649 .expect("relay service result");
1650 }
1651
1652 #[tokio::test]
1653 async fn regular_relay_bounds_close_reply_wait() {
1654 let (relay_ingress_dup, peer_ingress_dup) = duplex(kib(16));
1655 let (relay_egress_dup, peer_egress_dup) = duplex(kib(16));
1656
1657 let service = WebSocketRelayService::new(MirrorService::new())
1658 .with_close_handshake_timeout(Duration::from_millis(20));
1659 let relay = tokio::spawn(async move {
1660 service
1661 .serve(BridgeIo(
1662 MockSocket::new(relay_ingress_dup),
1663 MockSocket::new(relay_egress_dup),
1664 ))
1665 .await
1666 });
1667
1668 let mut peer_ingress_ws =
1669 AsyncWebSocket::from_raw_socket(MockSocket::new(peer_ingress_dup), Role::Client, None)
1670 .await;
1671 let mut peer_egress_ws =
1672 AsyncWebSocket::from_raw_socket(MockSocket::new(peer_egress_dup), Role::Server, None)
1673 .await;
1674
1675 let close_frame = test_close_frame("bounded shutdown");
1676 peer_ingress_ws
1677 .send_message(Message::Close(Some(close_frame.clone())))
1678 .await
1679 .expect("send ingress close");
1680 assert_eq!(
1681 expect_message(&mut peer_ingress_ws, "receive ingress close reply").await,
1682 Message::Close(Some(close_frame.clone()))
1683 );
1684 assert_eq!(
1685 expect_message(&mut peer_egress_ws, "receive propagated egress close").await,
1686 Message::Close(Some(close_frame))
1687 );
1688
1689 timeout(Duration::from_secs(1), relay)
1692 .await
1693 .expect("relay did not enforce its close handshake timeout")
1694 .expect("relay task join")
1695 .expect("relay service result");
1696 }
1697
1698 #[derive(Clone)]
1699 struct RecordingEventMiddleware {
1700 events: Arc<Mutex<Vec<(WebSocketRelayDirection, WebSocketRelayEvent)>>>,
1701 }
1702
1703 impl Service<WebSocketRelayEventInput> for RecordingEventMiddleware {
1704 type Output = WebSocketRelayEventOutput;
1705 type Error = BoxError;
1706
1707 async fn serve(
1708 &self,
1709 input: WebSocketRelayEventInput,
1710 ) -> Result<Self::Output, Self::Error> {
1711 self.events
1712 .lock()
1713 .push((input.direction, input.event.clone()));
1714 Ok(input.into())
1715 }
1716 }
1717
1718 #[derive(Clone)]
1719 struct FailingCloseObserver;
1720
1721 impl Service<WebSocketRelayEventInput> for FailingCloseObserver {
1722 type Output = WebSocketRelayEventOutput;
1723 type Error = BoxError;
1724
1725 async fn serve(
1726 &self,
1727 _input: WebSocketRelayEventInput,
1728 ) -> Result<Self::Output, Self::Error> {
1729 Err(BoxError::from_static_str("close observation failed"))
1730 }
1731 }
1732
1733 #[derive(Clone)]
1734 struct StallingCloseObserver;
1735
1736 impl Service<WebSocketRelayEventInput> for StallingCloseObserver {
1737 type Output = WebSocketRelayEventOutput;
1738 type Error = BoxError;
1739
1740 async fn serve(
1741 &self,
1742 input: WebSocketRelayEventInput,
1743 ) -> Result<Self::Output, Self::Error> {
1744 if matches!(&input.event, WebSocketRelayEvent::Close(_)) {
1745 return pending().await;
1746 }
1747 Ok(input.into())
1748 }
1749 }
1750
1751 #[tokio::test]
1752 async fn incoming_close_is_propagated_even_when_event_middleware_fails() {
1753 let (relay_ingress_dup, peer_ingress_dup) = duplex(kib(16));
1754 let (relay_egress_dup, peer_egress_dup) = duplex(kib(16));
1755
1756 let service = WebSocketRelayEventService::new(FailingCloseObserver);
1757 let relay = tokio::spawn(async move {
1758 service
1759 .serve(BridgeIo(
1760 MockSocket::new(relay_ingress_dup),
1761 MockSocket::new(relay_egress_dup),
1762 ))
1763 .await
1764 });
1765
1766 let mut peer_ingress_ws =
1767 AsyncWebSocket::from_raw_socket(MockSocket::new(peer_ingress_dup), Role::Client, None)
1768 .await;
1769 let mut peer_egress_ws =
1770 AsyncWebSocket::from_raw_socket(MockSocket::new(peer_egress_dup), Role::Server, None)
1771 .await;
1772
1773 let close_frame = test_close_frame("observer failure");
1774 peer_ingress_ws
1775 .send_message(Message::Close(Some(close_frame.clone())))
1776 .await
1777 .expect("send ingress close");
1778 assert_eq!(
1779 expect_message(&mut peer_ingress_ws, "receive ingress close reply").await,
1780 Message::Close(Some(close_frame.clone()))
1781 );
1782 assert_eq!(
1783 expect_message(&mut peer_egress_ws, "receive propagated egress close").await,
1784 Message::Close(Some(close_frame))
1785 );
1786 peer_egress_ws
1787 .flush()
1788 .await
1789 .expect("flush egress close reply");
1790 timeout(Duration::from_secs(1), relay)
1791 .await
1792 .expect("relay close timeout")
1793 .expect("relay task join")
1794 .expect("relay service result");
1795 }
1796
1797 #[tokio::test]
1798 async fn incoming_close_completion_is_not_delayed_by_stalled_observer() {
1799 let (relay_ingress_dup, peer_ingress_dup) = duplex(kib(16));
1800 let (relay_egress_dup, peer_egress_dup) = duplex(kib(16));
1801
1802 let service = WebSocketRelayEventService::new(StallingCloseObserver);
1803 let relay = tokio::spawn(async move {
1804 service
1805 .serve(BridgeIo(
1806 MockSocket::new(relay_ingress_dup),
1807 MockSocket::new(relay_egress_dup),
1808 ))
1809 .await
1810 });
1811
1812 let mut peer_ingress_ws =
1813 AsyncWebSocket::from_raw_socket(MockSocket::new(peer_ingress_dup), Role::Client, None)
1814 .await;
1815 let mut peer_egress_ws =
1816 AsyncWebSocket::from_raw_socket(MockSocket::new(peer_egress_dup), Role::Server, None)
1817 .await;
1818
1819 let close_frame = test_close_frame("observer stalls");
1820 peer_ingress_ws
1821 .send_message(Message::Close(Some(close_frame.clone())))
1822 .await
1823 .expect("send ingress close");
1824 assert_eq!(
1825 expect_message(&mut peer_ingress_ws, "receive ingress close reply").await,
1826 Message::Close(Some(close_frame.clone()))
1827 );
1828 assert_eq!(
1829 expect_message(&mut peer_egress_ws, "receive propagated egress close").await,
1830 Message::Close(Some(close_frame))
1831 );
1832 peer_egress_ws
1833 .flush()
1834 .await
1835 .expect("flush egress close reply");
1836
1837 timeout(Duration::from_secs(1), relay)
1838 .await
1839 .expect("stalled close observer delayed relay completion")
1840 .expect("relay task join")
1841 .expect("relay service result");
1842 }
1843
1844 #[tokio::test]
1845 async fn event_middleware_error_closes_both_connections_with_1011() {
1846 let (relay_ingress_dup, peer_ingress_dup) = duplex(kib(16));
1847 let (relay_egress_dup, peer_egress_dup) = duplex(kib(16));
1848
1849 let service = WebSocketRelayEventService::new(FailingCloseObserver);
1850 let relay = tokio::spawn(async move {
1851 service
1852 .serve(BridgeIo(
1853 MockSocket::new(relay_ingress_dup),
1854 MockSocket::new(relay_egress_dup),
1855 ))
1856 .await
1857 });
1858
1859 let mut peer_ingress_ws =
1860 AsyncWebSocket::from_raw_socket(MockSocket::new(peer_ingress_dup), Role::Client, None)
1861 .await;
1862 let mut peer_egress_ws =
1863 AsyncWebSocket::from_raw_socket(MockSocket::new(peer_egress_dup), Role::Server, None)
1864 .await;
1865
1866 peer_ingress_ws
1867 .send_message(Message::text("middleware fails"))
1868 .await
1869 .expect("send ingress text");
1870
1871 for (peer, description) in [
1872 (&mut peer_ingress_ws, "receive ingress error close"),
1873 (&mut peer_egress_ws, "receive egress error close"),
1874 ] {
1875 match expect_message(peer, description).await {
1876 Message::Close(Some(frame)) => assert_eq!(frame.code, CloseCode::Error),
1877 other => panic!("unexpected message while {description}: {other:?}"),
1878 }
1879 }
1880
1881 peer_ingress_ws
1882 .flush()
1883 .await
1884 .expect("flush ingress close reply");
1885 peer_egress_ws
1886 .flush()
1887 .await
1888 .expect("flush egress close reply");
1889 timeout(Duration::from_secs(1), relay)
1890 .await
1891 .expect("relay close timeout")
1892 .expect("relay task join")
1893 .expect("relay service result");
1894 }
1895
1896 #[tokio::test]
1897 async fn event_relay_observes_controls_without_forwarding_them() {
1898 let (relay_ingress_dup, peer_ingress_dup) = duplex(kib(16));
1899 let (relay_egress_dup, peer_egress_dup) = duplex(kib(16));
1900
1901 let events = Arc::new(Mutex::new(Vec::new()));
1902 let service = WebSocketRelayIoService::new(WebSocketRelayEventService::new(
1903 RecordingEventMiddleware {
1904 events: events.clone(),
1905 },
1906 ));
1907 let relay = tokio::spawn(async move {
1908 service
1909 .serve(BridgeIo(
1910 MockSocket::new(relay_ingress_dup),
1911 MockSocket::new(relay_egress_dup),
1912 ))
1913 .await
1914 });
1915
1916 let mut peer_ingress_ws =
1917 AsyncWebSocket::from_raw_socket(MockSocket::new(peer_ingress_dup), Role::Client, None)
1918 .await;
1919 let mut peer_egress_ws =
1920 AsyncWebSocket::from_raw_socket(MockSocket::new(peer_egress_dup), Role::Server, None)
1921 .await;
1922
1923 peer_ingress_ws
1924 .send_message(Message::text("hello"))
1925 .await
1926 .expect("send ingress text");
1927 assert_eq!(
1928 expect_message(&mut peer_egress_ws, "receive mirrored text").await,
1929 Message::text("hello")
1930 );
1931
1932 let ping = Bytes::from_static(b"observed-ping");
1933 peer_ingress_ws
1934 .send_message(Message::Ping(ping.clone()))
1935 .await
1936 .expect("send observed ping");
1937 assert_eq!(
1938 expect_message(
1939 &mut peer_ingress_ws,
1940 "receive observed ping's automatic pong",
1941 )
1942 .await,
1943 Message::Pong(ping.clone())
1944 );
1945 assert_eq!(
1946 expect_message(&mut peer_egress_ws, "receive observed ping heartbeat").await,
1947 Message::Pong(ping.clone())
1948 );
1949
1950 let pong = Bytes::from_static(b"observed-pong");
1951 peer_egress_ws
1952 .send_message(Message::Pong(pong.clone()))
1953 .await
1954 .expect("send observed pong");
1955 assert_no_message(&mut peer_ingress_ws, "ingress peer after observed pong").await;
1956
1957 let close_frame = test_close_frame("observed shutdown");
1958 peer_egress_ws
1959 .send_message(Message::Close(Some(close_frame.clone())))
1960 .await
1961 .expect("send observed close");
1962 assert_eq!(
1963 expect_message(&mut peer_egress_ws, "receive egress close reply").await,
1964 Message::Close(Some(close_frame.clone()))
1965 );
1966 assert_eq!(
1967 expect_message(&mut peer_ingress_ws, "receive propagated ingress close").await,
1968 Message::Close(Some(close_frame.clone()))
1969 );
1970 peer_ingress_ws
1971 .flush()
1972 .await
1973 .expect("flush ingress close reply");
1974 timeout(Duration::from_secs(1), relay)
1975 .await
1976 .expect("relay close timeout")
1977 .expect("relay task join")
1978 .expect("relay service result");
1979
1980 assert_eq!(
1981 *events.lock(),
1982 vec![
1983 (
1984 WebSocketRelayDirection::Ingress,
1985 WebSocketRelayEvent::Data(WebSocketRelayMessage::Text("hello".into())),
1986 ),
1987 (
1988 WebSocketRelayDirection::Ingress,
1989 WebSocketRelayEvent::Ping(ping),
1990 ),
1991 (
1992 WebSocketRelayDirection::Egress,
1993 WebSocketRelayEvent::Pong(pong),
1994 ),
1995 (
1996 WebSocketRelayDirection::Egress,
1997 WebSocketRelayEvent::Close(Some(close_frame)),
1998 ),
1999 ]
2000 );
2001 }
2002
2003 #[derive(Clone)]
2004 struct StallingIngressPingMiddleware {
2005 started: Arc<Mutex<Option<oneshot::Sender<()>>>>,
2006 }
2007
2008 impl Service<WebSocketRelayEventInput> for StallingIngressPingMiddleware {
2009 type Output = WebSocketRelayEventOutput;
2010 type Error = BoxError;
2011
2012 async fn serve(
2013 &self,
2014 input: WebSocketRelayEventInput,
2015 ) -> Result<Self::Output, Self::Error> {
2016 if input.direction == WebSocketRelayDirection::Ingress
2017 && matches!(&input.event, WebSocketRelayEvent::Ping(_))
2018 {
2019 if let Some(started) = self.started.lock().take() {
2020 _ = started.send(());
2021 }
2022 return pending().await;
2023 }
2024 Ok(input.into())
2025 }
2026 }
2027
2028 #[tokio::test]
2029 async fn stalled_middleware_does_not_block_the_opposite_direction_or_close() {
2030 let (relay_ingress_dup, peer_ingress_dup) = duplex(kib(16));
2031 let (relay_egress_dup, peer_egress_dup) = duplex(kib(16));
2032 let (started_tx, started_rx) = oneshot::channel();
2033
2034 let service = WebSocketRelayEventService::new(StallingIngressPingMiddleware {
2035 started: Arc::new(Mutex::new(Some(started_tx))),
2036 });
2037 let relay = tokio::spawn(async move {
2038 service
2039 .serve(BridgeIo(
2040 MockSocket::new(relay_ingress_dup),
2041 MockSocket::new(relay_egress_dup),
2042 ))
2043 .await
2044 });
2045
2046 let mut peer_ingress_ws =
2047 AsyncWebSocket::from_raw_socket(MockSocket::new(peer_ingress_dup), Role::Client, None)
2048 .await;
2049 let mut peer_egress_ws =
2050 AsyncWebSocket::from_raw_socket(MockSocket::new(peer_egress_dup), Role::Server, None)
2051 .await;
2052
2053 let ping = Bytes::from_static(b"stalling ping");
2054 peer_ingress_ws
2055 .send_message(Message::Ping(ping.clone()))
2056 .await
2057 .expect("send ingress ping");
2058 assert_eq!(
2059 expect_message(&mut peer_ingress_ws, "receive automatic ingress pong").await,
2060 Message::Pong(ping.clone())
2061 );
2062 timeout(Duration::from_secs(1), started_rx)
2063 .await
2064 .expect("middleware was not called")
2065 .expect("middleware start sender dropped");
2066 assert_eq!(
2067 expect_message(&mut peer_egress_ws, "receive egress heartbeat").await,
2068 Message::Pong(ping)
2069 );
2070
2071 peer_egress_ws
2072 .send_message(Message::text("opposite direction stays live"))
2073 .await
2074 .expect("send egress text");
2075 assert_eq!(
2076 expect_message(&mut peer_ingress_ws, "receive egress text").await,
2077 Message::text("opposite direction stays live")
2078 );
2079
2080 let close_frame = test_close_frame("opposite close");
2081 peer_egress_ws
2082 .send_message(Message::Close(Some(close_frame.clone())))
2083 .await
2084 .expect("send egress close");
2085 assert_eq!(
2086 expect_message(&mut peer_egress_ws, "receive egress close reply").await,
2087 Message::Close(Some(close_frame.clone()))
2088 );
2089 assert_eq!(
2090 expect_message(&mut peer_ingress_ws, "receive propagated ingress close").await,
2091 Message::Close(Some(close_frame))
2092 );
2093 peer_ingress_ws
2094 .flush()
2095 .await
2096 .expect("flush ingress close reply");
2097
2098 timeout(Duration::from_secs(1), relay)
2099 .await
2100 .expect("relay close timeout")
2101 .expect("relay task join")
2102 .expect("relay service result");
2103 }
2104
2105 #[derive(Clone)]
2106 struct CloseAfterDataMiddleware {
2107 close: WebSocketRelayClose,
2108 }
2109
2110 impl Service<WebSocketRelayEventInput> for CloseAfterDataMiddleware {
2111 type Output = WebSocketRelayEventOutput;
2112 type Error = BoxError;
2113
2114 async fn serve(
2115 &self,
2116 input: WebSocketRelayEventInput,
2117 ) -> Result<Self::Output, Self::Error> {
2118 let WebSocketRelayEventInput {
2119 direction: _,
2120 event,
2121 extensions,
2122 } = input;
2123 let messages = match event {
2124 WebSocketRelayEvent::Data(message) => vec![message],
2125 WebSocketRelayEvent::Ping(_)
2126 | WebSocketRelayEvent::Pong(_)
2127 | WebSocketRelayEvent::Close(_) => Vec::new(),
2128 };
2129 Ok(WebSocketRelayEventOutput {
2130 messages,
2131 close: Some(self.close.clone()),
2132 extensions,
2133 })
2134 }
2135 }
2136
2137 #[tokio::test]
2138 async fn event_relay_sends_messages_before_requested_coordinated_close() {
2139 let (relay_ingress_dup, peer_ingress_dup) = duplex(kib(16));
2140 let (relay_egress_dup, peer_egress_dup) = duplex(kib(16));
2141
2142 let close_frame = test_close_frame("middleware shutdown");
2143 let service = WebSocketRelayEventService::new(CloseAfterDataMiddleware {
2144 close: WebSocketRelayClose::WithFrame(close_frame.clone()),
2145 });
2146 let relay = tokio::spawn(async move {
2147 service
2148 .serve(BridgeIo(
2149 MockSocket::new(relay_ingress_dup),
2150 MockSocket::new(relay_egress_dup),
2151 ))
2152 .await
2153 });
2154
2155 let mut peer_ingress_ws =
2156 AsyncWebSocket::from_raw_socket(MockSocket::new(peer_ingress_dup), Role::Client, None)
2157 .await;
2158 let mut peer_egress_ws =
2159 AsyncWebSocket::from_raw_socket(MockSocket::new(peer_egress_dup), Role::Server, None)
2160 .await;
2161
2162 peer_egress_ws
2163 .send_message(Message::text("last message"))
2164 .await
2165 .expect("send final egress text");
2166 assert_eq!(
2167 expect_message(&mut peer_ingress_ws, "receive final ingress text").await,
2168 Message::text("last message")
2169 );
2170 assert_eq!(
2171 expect_message(&mut peer_egress_ws, "receive requested egress close").await,
2172 Message::Close(Some(close_frame.clone()))
2173 );
2174 assert_eq!(
2175 expect_message(&mut peer_ingress_ws, "receive requested ingress close").await,
2176 Message::Close(Some(close_frame))
2177 );
2178
2179 peer_ingress_ws
2180 .flush()
2181 .await
2182 .expect("flush ingress close reply");
2183 peer_egress_ws
2184 .flush()
2185 .await
2186 .expect("flush egress close reply");
2187 timeout(Duration::from_secs(1), relay)
2188 .await
2189 .expect("relay close timeout")
2190 .expect("relay task join")
2191 .expect("relay service result");
2192 }
2193
2194 #[tokio::test]
2195 async fn event_relay_can_request_frameless_close() {
2196 let (relay_ingress_dup, peer_ingress_dup) = duplex(kib(16));
2197 let (relay_egress_dup, peer_egress_dup) = duplex(kib(16));
2198
2199 let service = WebSocketRelayEventService::new(CloseAfterDataMiddleware {
2200 close: WebSocketRelayClose::WithoutFrame,
2201 });
2202 let relay = tokio::spawn(async move {
2203 service
2204 .serve(BridgeIo(
2205 MockSocket::new(relay_ingress_dup),
2206 MockSocket::new(relay_egress_dup),
2207 ))
2208 .await
2209 });
2210
2211 let mut peer_ingress_ws =
2212 AsyncWebSocket::from_raw_socket(MockSocket::new(peer_ingress_dup), Role::Client, None)
2213 .await;
2214 let mut peer_egress_ws =
2215 AsyncWebSocket::from_raw_socket(MockSocket::new(peer_egress_dup), Role::Server, None)
2216 .await;
2217
2218 peer_ingress_ws
2219 .send_message(Message::text("last frameless message"))
2220 .await
2221 .expect("send final ingress text");
2222 assert_eq!(
2223 expect_message(&mut peer_egress_ws, "receive final egress text").await,
2224 Message::text("last frameless message")
2225 );
2226 assert_eq!(
2227 expect_message(&mut peer_ingress_ws, "receive frameless ingress close").await,
2228 Message::Close(None)
2229 );
2230 assert_eq!(
2231 expect_message(&mut peer_egress_ws, "receive frameless egress close").await,
2232 Message::Close(None)
2233 );
2234
2235 peer_ingress_ws
2236 .flush()
2237 .await
2238 .expect("flush ingress close reply");
2239 peer_egress_ws
2240 .flush()
2241 .await
2242 .expect("flush egress close reply");
2243 timeout(Duration::from_secs(1), relay)
2244 .await
2245 .expect("relay close timeout")
2246 .expect("relay task join")
2247 .expect("relay service result");
2248 }
2249
2250 #[tokio::test]
2251 async fn event_relay_can_request_close_from_ping() {
2252 let (relay_ingress_dup, peer_ingress_dup) = duplex(kib(16));
2253 let (relay_egress_dup, peer_egress_dup) = duplex(kib(16));
2254
2255 let close_frame = test_close_frame("close on ping");
2256 let service = WebSocketRelayEventService::new(CloseAfterDataMiddleware {
2257 close: WebSocketRelayClose::WithFrame(close_frame.clone()),
2258 });
2259 let relay = tokio::spawn(async move {
2260 service
2261 .serve(BridgeIo(
2262 MockSocket::new(relay_ingress_dup),
2263 MockSocket::new(relay_egress_dup),
2264 ))
2265 .await
2266 });
2267
2268 let mut peer_ingress_ws =
2269 AsyncWebSocket::from_raw_socket(MockSocket::new(peer_ingress_dup), Role::Client, None)
2270 .await;
2271 let mut peer_egress_ws =
2272 AsyncWebSocket::from_raw_socket(MockSocket::new(peer_egress_dup), Role::Server, None)
2273 .await;
2274
2275 let ping = Bytes::from_static(b"close trigger");
2276 peer_ingress_ws
2277 .send_message(Message::Ping(ping.clone()))
2278 .await
2279 .expect("send ingress ping");
2280 assert_eq!(
2281 expect_message(&mut peer_ingress_ws, "receive automatic ingress pong").await,
2282 Message::Pong(ping.clone())
2283 );
2284 assert_eq!(
2285 expect_message(&mut peer_egress_ws, "receive egress heartbeat").await,
2286 Message::Pong(ping)
2287 );
2288 assert_eq!(
2289 expect_message(&mut peer_ingress_ws, "receive requested ingress close").await,
2290 Message::Close(Some(close_frame.clone()))
2291 );
2292 assert_eq!(
2293 expect_message(&mut peer_egress_ws, "receive requested egress close").await,
2294 Message::Close(Some(close_frame))
2295 );
2296
2297 peer_ingress_ws
2298 .flush()
2299 .await
2300 .expect("flush ingress close reply");
2301 peer_egress_ws
2302 .flush()
2303 .await
2304 .expect("flush egress close reply");
2305 timeout(Duration::from_secs(1), relay)
2306 .await
2307 .expect("relay close timeout")
2308 .expect("relay task join")
2309 .expect("relay service result");
2310 }
2311
2312 #[tokio::test]
2313 async fn event_relay_rejects_invalid_requested_close_frame() {
2314 let (relay_ingress_dup, peer_ingress_dup) = duplex(kib(16));
2315 let (relay_egress_dup, peer_egress_dup) = duplex(kib(16));
2316
2317 let invalid_close = CloseFrame {
2318 code: CloseCode::Normal,
2319 reason: "x".repeat(124).into(),
2320 };
2321 let service = WebSocketRelayEventService::new(CloseAfterDataMiddleware {
2322 close: WebSocketRelayClose::WithFrame(invalid_close),
2323 });
2324 let relay = tokio::spawn(async move {
2325 service
2326 .serve(BridgeIo(
2327 MockSocket::new(relay_ingress_dup),
2328 MockSocket::new(relay_egress_dup),
2329 ))
2330 .await
2331 });
2332
2333 let mut peer_ingress_ws =
2334 AsyncWebSocket::from_raw_socket(MockSocket::new(peer_ingress_dup), Role::Client, None)
2335 .await;
2336 let mut peer_egress_ws =
2337 AsyncWebSocket::from_raw_socket(MockSocket::new(peer_egress_dup), Role::Server, None)
2338 .await;
2339
2340 peer_ingress_ws
2341 .send_message(Message::text("must not be relayed"))
2342 .await
2343 .expect("send ingress text");
2344 for (peer, description) in [
2345 (&mut peer_ingress_ws, "receive ingress invalid-output close"),
2346 (&mut peer_egress_ws, "receive egress invalid-output close"),
2347 ] {
2348 match expect_message(peer, description).await {
2349 Message::Close(Some(frame)) => assert_eq!(frame.code, CloseCode::Error),
2350 other => panic!("unexpected message while {description}: {other:?}"),
2351 }
2352 }
2353
2354 peer_ingress_ws
2355 .flush()
2356 .await
2357 .expect("flush ingress close reply");
2358 peer_egress_ws
2359 .flush()
2360 .await
2361 .expect("flush egress close reply");
2362 timeout(Duration::from_secs(1), relay)
2363 .await
2364 .expect("relay close timeout")
2365 .expect("relay task join")
2366 .expect("relay service result");
2367 }
2368
2369 #[test]
2370 fn event_output_mirrors_only_data_and_distinguishes_empty_close() {
2371 let extensions = Extensions::new();
2372 extensions.insert(IngressMarker);
2373 let input = WebSocketRelayInput {
2374 direction: WebSocketRelayDirection::Ingress,
2375 message: WebSocketRelayMessage::Text("input".into()),
2376 extensions,
2377 };
2378 assert!(input.extensions().contains::<IngressMarker>());
2379
2380 let extensions = Extensions::new();
2381 extensions.insert(IngressMarker);
2382 let output = WebSocketRelayOutput {
2383 messages: Vec::new(),
2384 extensions,
2385 };
2386 assert!(output.extensions().contains::<IngressMarker>());
2387
2388 let extensions = Extensions::new();
2389 extensions.insert(IngressMarker);
2390 let event_input = WebSocketRelayEventInput {
2391 direction: WebSocketRelayDirection::Ingress,
2392 event: WebSocketRelayEvent::Ping(Bytes::new()),
2393 extensions,
2394 };
2395 assert!(event_input.extensions().contains::<IngressMarker>());
2396
2397 let extensions = Extensions::new();
2398 extensions.insert(IngressMarker);
2399 let event_output = WebSocketRelayEventOutput {
2400 messages: Vec::new(),
2401 close: None,
2402 extensions,
2403 };
2404 assert!(event_output.extensions().contains::<IngressMarker>());
2405
2406 let data_output: WebSocketRelayEventOutput = WebSocketRelayEventInput {
2407 direction: WebSocketRelayDirection::Ingress,
2408 event: WebSocketRelayEvent::Data(WebSocketRelayMessage::Binary(Bytes::from_static(
2409 b"data",
2410 ))),
2411 extensions: Extensions::new(),
2412 }
2413 .into();
2414 assert_eq!(
2415 data_output.messages,
2416 vec![WebSocketRelayMessage::Binary(Bytes::from_static(b"data"))]
2417 );
2418 assert_eq!(data_output.close, None);
2419
2420 for event in [
2421 WebSocketRelayEvent::Ping(Bytes::from_static(b"ping")),
2422 WebSocketRelayEvent::Pong(Bytes::from_static(b"pong")),
2423 WebSocketRelayEvent::Close(None),
2424 ] {
2425 let output: WebSocketRelayEventOutput = WebSocketRelayEventInput {
2426 direction: WebSocketRelayDirection::Egress,
2427 event,
2428 extensions: Extensions::new(),
2429 }
2430 .into();
2431 assert!(output.messages.is_empty());
2432 assert_eq!(output.close, None);
2433 }
2434
2435 assert_eq!(WebSocketRelayClose::WithoutFrame.into_frame(), None);
2436 let close_frame = test_close_frame("with frame");
2437 assert_eq!(
2438 WebSocketRelayClose::WithFrame(close_frame.clone()).into_frame(),
2439 Some(close_frame.clone())
2440 );
2441 assert_eq!(
2442 WebSocketRelayClose::from(Some(close_frame.clone())),
2443 WebSocketRelayClose::WithFrame(close_frame)
2444 );
2445 assert_eq!(
2446 WebSocketRelayClose::from(None),
2447 WebSocketRelayClose::WithoutFrame
2448 );
2449
2450 assert!(valid_close_frame(None));
2451 assert!(valid_close_frame(Some(&CloseFrame {
2452 code: CloseCode::Normal,
2453 reason: "x".repeat(123).into(),
2454 })));
2455 assert!(!valid_close_frame(Some(&CloseFrame {
2456 code: CloseCode::Normal,
2457 reason: "x".repeat(124).into(),
2458 })));
2459 assert!(!valid_close_frame(Some(&CloseFrame {
2460 code: CloseCode::Abnormal,
2461 reason: "forbidden code".into(),
2462 })));
2463 }
2464
2465 #[test]
2466 fn relay_services_expose_close_handshake_timeout_setters() {
2467 let timeout = Duration::from_secs(7);
2468
2469 let mut service = WebSocketRelayService::new(MirrorService::new());
2470 service.set_close_handshake_timeout(timeout);
2471 assert_eq!(service.close_handshake_timeout, timeout);
2472
2473 let service = WebSocketRelayEventService::new(MirrorService::new())
2474 .with_close_handshake_timeout(timeout);
2475 assert_eq!(service.close_handshake_timeout, timeout);
2476 }
2477}