Skip to main content

rama_ws/handshake/
mitm.rs

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/// A pair of established WebSocket message transports joined by a relay.
26///
27/// The ingress side faces the downstream client and uses the server protocol
28/// role. The egress side faces the upstream server and uses the client role.
29/// Keeping this boundary distinct from [`BridgeIo`] lets ordinary Rama layers
30/// decorate complete WebSocket messages without entering the protocol runtime
31/// or decoding raw bytes themselves. The bridge deliberately does not select
32/// one side as its extension source; use [`Self::ingress`] or [`Self::egress`]
33/// to access the intended transport explicitly.
34#[derive(Debug)]
35pub struct WebSocketBridge<Ingress, Egress> {
36    /// Message transport facing the downstream client.
37    pub ingress: Ingress,
38    /// Message transport facing the upstream server.
39    pub egress: Egress,
40}
41
42/// Adapt a raw byte-level [`BridgeIo`] into a [`WebSocketBridge`] before
43/// invoking an inner service.
44///
45/// This adapter is useful when message-level layers must sit between protocol
46/// construction and a relay. Existing [`WebSocketRelayService`] and
47/// [`WebSocketRelayEventService`] values continue to accept raw [`BridgeIo`]
48/// directly for backwards compatibility.
49#[derive(Debug, Clone)]
50pub struct WebSocketRelayIoService<S> {
51    inner: S,
52}
53
54/// Layer that adapts a raw byte-level [`BridgeIo`] into a message-level
55/// [`WebSocketBridge`].
56#[derive(Debug, Clone, Copy, Default)]
57pub struct WebSocketRelayIoLayer;
58
59impl WebSocketRelayIoLayer {
60    /// Create a WebSocket relay I/O adapter layer.
61    #[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    /// Create a raw-I/O adapter for a message-level WebSocket service.
77    #[must_use]
78    pub const fn new(inner: S) -> Self {
79        Self { inner }
80    }
81
82    /// Return a reference to the inner message-level service.
83    #[must_use]
84    pub const fn inner(&self) -> &S {
85        &self.inner
86    }
87
88    /// Consume this adapter and return its inner service.
89    #[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)]
112/// A utility that can be used by MITM services such as transparent proxies,
113/// in order to relay WebSocket messages.
114///
115/// By default they get mirrored but the logic is fully up to you.
116///
117/// This service accepts both a raw [`BridgeIo`] and an established
118/// [`WebSocketBridge`]. Direct raw-I/O use remains convenient and backwards
119/// compatible. To install message-level layers, wrap this service in those
120/// layers and then place [`WebSocketRelayIoService`] around the result.
121///
122/// ## KISS
123///
124/// This service is for simple DPI purposes.
125///
126/// Ping and pong are handled locally on each of the two independent WebSocket
127/// connections and are not exposed to middleware. A ping is acknowledged on
128/// its source connection and also produces an unsolicited pong heartbeat on
129/// the opposite connection, so activity on one leg keeps both legs alive
130/// without coupling their ping round trips. Other pongs are not forwarded. An
131/// incoming close starts coordinated shutdown; data received while closing is
132/// discarded rather than passed to middleware.
133///
134/// Middleware is processed independently per direction. Its future can be
135/// cancelled when either peer starts closing, so it must be cancel-safe. A
136/// failure to send middleware-produced data terminates the relay.
137///
138/// Use [`WebSocketRelayEventService`] when middleware also needs to observe
139/// control messages. Fork or create your own relay service for lower-level
140/// purposes such as preserving raw frame boundaries.
141pub 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    /// Create a new [`WebSocketRelayService`]
150    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        /// Set how long the relay waits for both peers to finish a coordinated
159        /// close handshake before dropping the connections.
160        ///
161        /// The default is five seconds. Both connections and their relay state
162        /// remain alive until the handshake finishes or this timeout expires.
163        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)]
171/// A WebSocket MITM relay that exposes every message observable through the
172/// high-level WebSocket protocol API to its middleware.
173///
174/// Like [`WebSocketRelayService`], this accepts either raw [`BridgeIo`] or an
175/// established [`WebSocketBridge`]. Use [`WebSocketRelayIoService`] when
176/// message-level layers must run between protocol construction and this relay.
177///
178/// Unlike [`WebSocketRelayService`], this service exposes ping, pong and close
179/// events. Control messages remain owned by the relay: a ping is acknowledged
180/// locally and produces an unsolicited pong heartbeat on the opposite
181/// connection, while other pongs are not forwarded. An incoming close always
182/// starts coordinated shutdown.
183///
184/// Middleware is processed independently per direction. Its future can be
185/// cancelled when either peer starts closing, so it must be cancel-safe. Data
186/// received after shutdown starts is discarded rather than exposed, and a
187/// failure to send middleware-produced data terminates the relay.
188pub 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    /// Create a new [`WebSocketRelayEventService`].
197    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        /// Set how long the relay waits for both peers to finish a coordinated
206        /// close handshake before dropping the connections.
207        ///
208        /// The default is five seconds. Both connections and their relay state
209        /// remain alive until the handshake finishes or this timeout expires.
210        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)]
218/// Most typically used as Input
219/// for users of [`WebSocketRelayService`].
220pub 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)]
234/// Most typically used as Output
235/// for users of [`WebSocketRelayService`].
236pub struct WebSocketRelayOutput {
237    /// 0 or more messages, providing the ability
238    /// to drop messages first and return buffered messages later.
239    /// Messages are sent to the opposite WebSocket connection.
240    pub messages: Vec<WebSocketRelayMessage>,
241    /// Per-direction relay state. Middleware should normally return the input
242    /// store (or a derivative of it); replacing it with a fresh store also
243    /// replaces its connection-extension parent link.
244    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)]
270/// Input for middleware used by [`WebSocketRelayEventService`].
271pub 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)]
285/// Output for middleware used by [`WebSocketRelayEventService`].
286pub struct WebSocketRelayEventOutput {
287    /// Zero or more data messages to send to the opposite WebSocket connection.
288    ///
289    /// Ping, pong and raw frames cannot be produced through this API.
290    pub messages: Vec<WebSocketRelayMessage>,
291    /// Optionally request coordinated shutdown of both WebSocket connections.
292    /// Messages are sent before a valid requested shutdown starts.
293    ///
294    /// When serving [`WebSocketRelayEvent::Close`], shutdown has already been
295    /// initiated by the relay and both `messages` and `close` are ignored.
296    pub close: Option<WebSocketRelayClose>,
297    /// Per-direction relay state. Middleware should normally return the input
298    /// store (or a derivative of it); replacing it with a fresh store also
299    /// replaces its connection-extension parent link.
300    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)]
334/// A message observed by [`WebSocketRelayEventService`] middleware.
335///
336/// Raw [`crate::protocol::frame::Frame`] values are intentionally absent:
337/// [`crate::Message::Frame`] is send-only and is never returned while reading.
338pub enum WebSocketRelayEvent {
339    /// An application data message.
340    Data(WebSocketRelayMessage),
341    /// A ping received from one WebSocket peer.
342    Ping(Bytes),
343    /// A pong received from one WebSocket peer.
344    Pong(Bytes),
345    /// A close received from one WebSocket peer.
346    Close(Option<CloseFrame>),
347}
348
349#[derive(Debug, Clone, Eq, PartialEq)]
350/// A coordinated close requested by [`WebSocketRelayEventOutput`].
351///
352/// This enum distinguishes no close request (`None`) from a close message that
353/// intentionally carries no status code or reason ([`Self::WithoutFrame`]).
354pub enum WebSocketRelayClose {
355    /// Close without a status code or reason.
356    WithoutFrame,
357    /// Close with a status code and optional reason.
358    ///
359    /// The relay rejects codes that cannot appear on the wire and reasons
360    /// longer than 123 bytes. The entire middleware output is then rejected:
361    /// its messages are discarded and both connections close with status 1011.
362    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)]
384/// Non-meta WebSocket messages, used as part of [`WebSocketRelayInput`]
385/// and [`WebSocketRelayOutput`], most typically for users of [`WebSocketRelayService`].
386pub enum WebSocketRelayMessage {
387    /// A text WebSocket message
388    Text(Utf8Bytes),
389    /// A binary WebSocket message
390    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)]
403/// Direction data used as part of [`WebSocketRelayInput`],
404/// most typically for users of [`WebSocketRelayService`].
405pub 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    // Each direction gets a child store of the socket the event arrived on.
663    // Middleware can see the live socket's extensions without mutating it or
664    // leaking inserts into the other relay direction.
665    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            // RFC 6455 permits unsolicited Pong frames as a unidirectional
1013            // heartbeat. This keeps the other independent relay leg active
1014            // without making the source peer wait on its round trip.
1015            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    //! End-to-end regression coverage for data/control routing, coordinated
1251    //! close behavior and per-direction middleware-extension isolation.
1252    //! The isolation test distinguishes a shared `clone()` (cross-direction
1253    //! marker leak) from per-direction `clone()` (live-socket pollution).
1254
1255    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                // Parent visibility: a fork() walks into the parent on lookup,
1352                // so the side's pre-inserted marker MUST be reachable.
1353                saw_ingress_marker: extensions.get_ref::<IngressMarker>().is_some(),
1354                saw_egress_marker: extensions.get_ref::<EgressMarker>().is_some(),
1355                // Cross-direction visibility: forks are independent, so neither
1356                // direction's middleware insert should be visible in the other.
1357                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        // Two duplex pairs: one for the ingress side of the relay, one for
1379        // the egress side. `MockSocket` wraps each end in an `ExtensionsRef`
1380        // shell so the relay's `from_raw_socket` is happy.
1381        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        // Capture handles to the live socket extension stores BEFORE
1390        // moving the sockets into the relay. `Extensions::clone()` shares
1391        // the top-level `Arc`, so any insert that ends up on the live
1392        // store would be observable through these handles after the
1393        // relay finishes. `fork()` does NOT share that `Arc`, so
1394        // correctly-forked inserts won't be observable here.
1395        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        // Relay's ingress is `Role::Server`, so the peer plays `Role::Client`
1406        // (masked frames). Egress is the mirror.
1407        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        // ingress -> egress
1415        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        // egress -> ingress
1425        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        // Dropping a peer closes its duplex end; the relay sees a connection
1435        // error and returns.
1436        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        // Per-direction parent visibility: fork() preserves walk-into-parent.
1453        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        // Cross-direction probe isolation. If the wiring regresses to a
1468        // single shared `egress_socket.extensions().clone()` threaded to
1469        // BOTH directions, the egress middleware call would see
1470        // `LeakProbeIngress` (and/or vice versa).
1471        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        // Live-socket isolation. If the wiring regresses to per-direction
1481        // `clone()`, the top-level `Arc` would be shared with the live
1482        // socket store, so the middleware insert would surface here.
1483        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        // Intentionally do not flush the egress peer's automatically queued
1690        // reply. The relay must still terminate at its configured bound.
1691        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}