Skip to main content

rama_ws/layer/
har.rs

1//! HAR capture middleware for WebSocket message streams.
2//!
3//! This module deliberately decorates an already constructed WebSocket. The
4//! core WebSocket runtime and protocol implementation do not know about HAR,
5//! recorders, or HTTP middleware extensions.
6//!
7//! [`HARWebSocketLayer`] is explicitly installed around a service that accepts
8//! a client endpoint, server endpoint, or established relay bridge. Manual
9//! handshake users can instead apply [`ClientWebSocket::map_socket`] or
10//! [`ServerWebSocket::map_socket`] and construct [`HARWebSocket`] directly.
11//! Code that does not apply either form keeps using the ordinary WebSocket
12//! types and pays no capture-wrapper cost.
13//!
14//! A relay capture records the messages finally accepted by its destination
15//! legs: writes to the upstream-facing egress socket are HAR `send` messages,
16//! while writes to the downstream-facing ingress socket are HAR `receive`
17//! messages. Messages dropped by relay middleware are therefore absent, and
18//! transformed or expanded outputs are represented as actually forwarded.
19//! The opaque capture handle is carried by the egress transport because it
20//! continues the HAR session created around the upstream HTTP handshake. Once
21//! claimed, that one session is shared by the wrappers on both bridge legs;
22//! its egress location does not limit observation to egress traffic.
23//!
24//! In an endpoint stack, apply this layer to the service that consumes a
25//! [`ClientWebSocket`] or [`ServerWebSocket`]. The layer reads the opaque
26//! capture handle from that wrapper's preserved HTTP handshake metadata,
27//! replaces only its generic socket parameter, and calls the inner endpoint.
28//! For a relay, place the same layer between
29//! [`WebSocketRelayIoService`](crate::handshake::mitm::WebSocketRelayIoService)
30//! and a message-level relay service:
31//!
32//! ```text
33//! BridgeIo<raw ingress, raw egress>
34//!     -> WebSocketRelayIoService
35//!     -> HARWebSocketLayer
36//!     -> WebSocketRelayService
37//! ```
38//!
39//! A manual client can clone [`WebSocketCapture`] from
40//! `websocket.response().extensions` and pass it to [`HARWebSocket::new`]
41//! through [`ClientWebSocket::map_socket`]. The equivalent server metadata is
42//! available through `websocket.request().extensions`. No HAR-specific
43//! extension trait or handshake variant is required.
44//!
45//! [`ClientWebSocket::map_socket`]: crate::handshake::client::ClientWebSocket::map_socket
46//! [`ServerWebSocket::map_socket`]: crate::handshake::server::ServerWebSocket::map_socket
47
48use crate::{
49    Message, ProtocolError, WebSocketIo,
50    handshake::{client::ClientWebSocket, mitm::WebSocketBridge, server::ServerWebSocket},
51    protocol::Role,
52};
53use rama_core::{
54    Layer, Service,
55    extensions::{Extensions, ExtensionsRef},
56    futures::{Sink, SinkExt as _, Stream, StreamExt as _, task::AtomicWaker},
57    telemetry::tracing::debug,
58};
59use rama_http::layer::har::{
60    recorder::{WebSocketCapture, WebSocketCaptureFuture, WebSocketCaptureLease},
61    spec::{WebSocketMessage, WebSocketMessageType},
62};
63use rama_utils::time::unix_timestamp_millis;
64use std::{
65    fmt,
66    future::Future,
67    io,
68    pin::Pin,
69    sync::Arc,
70    task::{Context, Poll, Wake, Waker, ready},
71};
72
73struct PendingObservation {
74    future: WebSocketCaptureFuture,
75    close_after: bool,
76}
77
78#[derive(Clone, Copy)]
79enum ObservationSide {
80    Read,
81    Write,
82}
83
84// `StreamExt::split` can poll one observation from two different tasks. Poll
85// the recorder with this fan-out waker so its latest registration always wakes
86// both halves, even if the half that polled last is subsequently cancelled.
87struct ObservationWakers {
88    read: AtomicWaker,
89    write: AtomicWaker,
90}
91
92impl ObservationWakers {
93    fn new() -> Self {
94        Self {
95            read: AtomicWaker::new(),
96            write: AtomicWaker::new(),
97        }
98    }
99
100    fn register(&self, side: ObservationSide, waker: &Waker) {
101        match side {
102            ObservationSide::Read => self.read.register(waker),
103            ObservationSide::Write => self.write.register(waker),
104        }
105    }
106
107    fn wake_waiters(&self) {
108        self.read.wake();
109        self.write.wake();
110    }
111}
112
113impl Wake for ObservationWakers {
114    fn wake(self: Arc<Self>) {
115        self.wake_waiters();
116    }
117
118    fn wake_by_ref(self: &Arc<Self>) {
119        self.wake_waiters();
120    }
121}
122
123/// Rama layer that installs HAR capture around WebSocket endpoint services or
124/// an established relay bridge.
125#[derive(Debug, Clone, Copy, Default)]
126pub struct HARWebSocketLayer;
127
128impl HARWebSocketLayer {
129    /// Create a WebSocket HAR service layer.
130    #[must_use]
131    pub const fn new() -> Self {
132        Self
133    }
134}
135
136impl<S> Layer<S> for HARWebSocketLayer {
137    type Service = HARWebSocketService<S>;
138
139    fn layer(&self, inner: S) -> Self::Service {
140        HARWebSocketService { inner }
141    }
142
143    fn into_layer(self, inner: S) -> Self::Service {
144        HARWebSocketService { inner }
145    }
146}
147
148/// Service produced by [`HARWebSocketLayer`].
149#[derive(Debug, Clone)]
150pub struct HARWebSocketService<S> {
151    inner: S,
152}
153
154impl<Inner, Socket> Service<ClientWebSocket<Socket>> for HARWebSocketService<Inner>
155where
156    Inner: Service<ClientWebSocket<HARWebSocket<Socket>>>,
157    Socket: WebSocketIo,
158{
159    type Output = Inner::Output;
160    type Error = Inner::Error;
161
162    async fn serve(&self, websocket: ClientWebSocket<Socket>) -> Result<Self::Output, Self::Error> {
163        let capture = websocket
164            .response()
165            .extensions
166            .get_ref::<WebSocketCapture>()
167            .cloned();
168        self.inner
169            .serve(
170                websocket
171                    .map_socket(move |socket| HARWebSocket::new(socket, Role::Client, capture)),
172            )
173            .await
174    }
175}
176
177impl<Inner, Socket> Service<ServerWebSocket<Socket>> for HARWebSocketService<Inner>
178where
179    Inner: Service<ServerWebSocket<HARWebSocket<Socket>>>,
180    Socket: WebSocketIo,
181{
182    type Output = Inner::Output;
183    type Error = Inner::Error;
184
185    async fn serve(&self, websocket: ServerWebSocket<Socket>) -> Result<Self::Output, Self::Error> {
186        let capture = websocket
187            .request()
188            .extensions
189            .get_ref::<WebSocketCapture>()
190            .cloned();
191        self.inner
192            .serve(
193                websocket
194                    .map_socket(move |socket| HARWebSocket::new(socket, Role::Server, capture)),
195            )
196            .await
197    }
198}
199
200impl<Inner, Ingress, Egress> Service<WebSocketBridge<Ingress, Egress>>
201    for HARWebSocketService<Inner>
202where
203    Inner: Service<WebSocketBridge<HARWebSocket<Ingress>, HARWebSocket<Egress>>>,
204    Ingress: WebSocketIo,
205    Egress: WebSocketIo,
206{
207    type Output = Inner::Output;
208    type Error = Inner::Error;
209
210    async fn serve(
211        &self,
212        WebSocketBridge { ingress, egress }: WebSocketBridge<Ingress, Egress>,
213    ) -> Result<Self::Output, Self::Error> {
214        // HARExportLayer creates this session around the upstream HTTP
215        // handshake. A successful response carries its continuation into the
216        // upgraded egress transport; the lease itself then observes writes on
217        // both legs of the message bridge.
218        let capture_lease = egress
219            .extensions()
220            .get_ref::<WebSocketCapture>()
221            .and_then(WebSocketCapture::lease)
222            .map(Arc::new);
223
224        let ingress = HARWebSocket::relay_leg(
225            ingress,
226            WebSocketMessageType::Receive,
227            capture_lease.clone(),
228        );
229        let egress =
230            HARWebSocket::relay_leg(egress, WebSocketMessageType::Send, capture_lease.clone());
231
232        self.inner.serve(WebSocketBridge { ingress, egress }).await
233    }
234}
235
236#[derive(Debug, Clone, Copy)]
237enum CaptureMode {
238    Endpoint(Role),
239    Writes(WebSocketMessageType),
240}
241
242impl CaptureMode {
243    fn message_type(self, outgoing: bool) -> Option<WebSocketMessageType> {
244        match (self, outgoing) {
245            (Self::Endpoint(Role::Client), true) | (Self::Endpoint(Role::Server), false) => {
246                Some(WebSocketMessageType::Send)
247            }
248            (Self::Endpoint(Role::Client), false) | (Self::Endpoint(Role::Server), true) => {
249                Some(WebSocketMessageType::Receive)
250            }
251            (Self::Writes(message_type), true) => Some(message_type),
252            (Self::Writes(_), false) => None,
253        }
254    }
255}
256
257impl fmt::Debug for PendingObservation {
258    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
259        formatter
260            .debug_struct("PendingObservation")
261            .field("close_after", &self.close_after)
262            .finish_non_exhaustive()
263    }
264}
265
266/// HAR-capturing middleware around a WebSocket message stream.
267///
268/// The wrapper observes complete text and binary messages after the underlying
269/// WebSocket protocol accepts or produces them. It awaits the configured
270/// recorder before accepting the next socket operation, which applies bounded
271/// backpressure without adding capture concerns to the underlying WebSocket.
272///
273/// This type intentionally does not dereference to `S`, because directly
274/// polling the inner transport would bypass capture. Use [`Self::get_ref`] or
275/// [`Self::get_mut`] only when that bypass is explicitly intended, and call
276/// [`Self::into_inner`] to remove the middleware.
277pub struct HARWebSocket<S> {
278    inner: S,
279    mode: CaptureMode,
280    capture_lease: Option<Arc<WebSocketCaptureLease>>,
281    close_on_terminal: bool,
282    pending_observation: Option<PendingObservation>,
283    queued_observation: Option<PendingObservation>,
284    observation_wakers: Option<Arc<ObservationWakers>>,
285    pending_read: Option<Result<Message, ProtocolError>>,
286    pending_write_error: Option<ProtocolError>,
287}
288
289impl<S: fmt::Debug> fmt::Debug for HARWebSocket<S> {
290    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
291        formatter
292            .debug_struct("HARWebSocket")
293            .field("inner", &self.inner)
294            .field("mode", &self.mode)
295            .field("capture_lease", &self.capture_lease)
296            .field("pending_observation", &self.pending_observation)
297            .field("queued_observation", &self.queued_observation)
298            .field("pending_read", &self.pending_read)
299            .field("pending_write_error", &self.pending_write_error)
300            .finish()
301    }
302}
303
304impl<S> HARWebSocket<S> {
305    /// Wrap a WebSocket with an optional opaque capture handle.
306    #[must_use]
307    pub fn new(inner: S, role: Role, capture: Option<WebSocketCapture>) -> Self {
308        Self::from_parts(
309            inner,
310            CaptureMode::Endpoint(role),
311            capture.and_then(|capture| capture.lease()).map(Arc::new),
312            true,
313        )
314    }
315
316    fn relay_leg(
317        inner: S,
318        message_type: WebSocketMessageType,
319        capture_lease: Option<Arc<WebSocketCaptureLease>>,
320    ) -> Self {
321        Self::from_parts(
322            inner,
323            CaptureMode::Writes(message_type),
324            capture_lease,
325            false,
326        )
327    }
328
329    fn from_parts(
330        inner: S,
331        mode: CaptureMode,
332        capture_lease: Option<Arc<WebSocketCaptureLease>>,
333        close_on_terminal: bool,
334    ) -> Self {
335        let observation_wakers = capture_lease
336            .as_ref()
337            .map(|_| Arc::new(ObservationWakers::new()));
338        Self {
339            inner,
340            mode,
341            capture_lease,
342            close_on_terminal,
343            pending_observation: None,
344            queued_observation: None,
345            observation_wakers,
346            pending_read: None,
347            pending_write_error: None,
348        }
349    }
350
351    /// Wrap a WebSocket using the capture handle reachable from its extensions.
352    #[must_use]
353    pub fn from_extensions(inner: S, role: Role) -> Self
354    where
355        S: ExtensionsRef,
356    {
357        let capture = inner.extensions().get_ref::<WebSocketCapture>().cloned();
358        Self::new(inner, role, capture)
359    }
360
361    /// Remove this middleware and return the underlying WebSocket.
362    #[must_use]
363    pub fn into_inner(self) -> S {
364        self.inner
365    }
366
367    /// Return a shared reference to the underlying WebSocket.
368    #[must_use]
369    pub fn get_ref(&self) -> &S {
370        &self.inner
371    }
372
373    /// Return a mutable reference to the underlying WebSocket.
374    #[must_use]
375    pub fn get_mut(&mut self) -> &mut S {
376        &mut self.inner
377    }
378
379    fn poll_observation(&mut self, ctx: &Context<'_>, side: ObservationSide) -> Poll<()> {
380        loop {
381            let Some(observation) = &mut self.pending_observation else {
382                return Poll::Ready(());
383            };
384            let Some(observation_wakers) = self.observation_wakers.as_ref() else {
385                self.pending_observation.take();
386                self.queued_observation.take();
387                return Poll::Ready(());
388            };
389            observation_wakers.register(side, ctx.waker());
390            let observation_waker = Waker::from(observation_wakers.clone());
391            let mut observation_ctx = Context::from_waker(&observation_waker);
392            match Pin::new(&mut observation.future).poll(&mut observation_ctx) {
393                Poll::Pending => return Poll::Pending,
394                Poll::Ready(result) => {
395                    observation_wakers.wake_waiters();
396                    let close_after = self
397                        .pending_observation
398                        .take()
399                        .is_some_and(|observation| observation.close_after);
400                    if let Err(err) = &result {
401                        debug!("failed to record WebSocket HAR observation: {err}");
402                    }
403                    if result.is_err() || close_after {
404                        if let Some(capture) = &self.capture_lease {
405                            capture.close();
406                        }
407                        self.capture_lease.take();
408                        self.queued_observation.take();
409                        return Poll::Ready(());
410                    }
411                    self.pending_observation = self.queued_observation.take();
412                }
413            }
414        }
415    }
416
417    fn queue_observation(&mut self, observation: PendingObservation) {
418        if self.observation_wakers.is_none() {
419            self.observation_wakers = Some(Arc::new(ObservationWakers::new()));
420        }
421        if self.pending_observation.is_none() {
422            self.pending_observation = Some(observation);
423        } else if self.queued_observation.is_none() {
424            self.queued_observation = Some(observation);
425        } else {
426            debug!("discarding WebSocket HAR observation after Sink contract violation");
427            debug_assert!(
428                false,
429                "calling start_send repeatedly without poll_ready violates Sink"
430            );
431        }
432    }
433
434    fn message_observation(
435        &self,
436        outgoing: bool,
437        message: &Message,
438    ) -> Option<WebSocketCaptureFuture> {
439        let capture = self.capture_lease.as_ref()?;
440        if capture.is_closed() {
441            return None;
442        }
443        let message_type = self.mode.message_type(outgoing)?;
444        into_har_message(message_type, message).map(|message| capture.record(message))
445    }
446
447    fn begin_message_observation(
448        &mut self,
449        outgoing: bool,
450        message: &Message,
451        close_after: bool,
452    ) -> bool {
453        if let Some(future) = self.message_observation(outgoing, message) {
454            self.queue_observation(PendingObservation {
455                future,
456                close_after,
457            });
458            true
459        } else {
460            if close_after && self.close_on_terminal {
461                if let Some(capture) = &self.capture_lease {
462                    capture.close();
463                }
464                self.capture_lease.take();
465            }
466            false
467        }
468    }
469
470    fn begin_error_observation(&mut self, error: &ProtocolError) -> bool {
471        let Some(capture) = &self.capture_lease else {
472            return false;
473        };
474        if capture.is_closed() {
475            return false;
476        }
477        let future = capture.record(WebSocketMessage::error(
478            epoch_seconds_from_millis(unix_timestamp_millis()),
479            error.to_string(),
480        ));
481        self.queue_observation(PendingObservation {
482            future,
483            close_after: self.close_on_terminal,
484        });
485        true
486    }
487
488    fn set_pending_observation(&mut self, future: WebSocketCaptureFuture) {
489        self.queue_observation(PendingObservation {
490            future,
491            close_after: false,
492        });
493    }
494
495    fn poll_write_error(
496        &mut self,
497        ctx: &Context<'_>,
498        error: ProtocolError,
499    ) -> Poll<Result<(), ProtocolError>> {
500        if !self.begin_error_observation(&error) {
501            return Poll::Ready(Err(error));
502        }
503        self.pending_write_error = Some(error);
504        if self
505            .poll_observation(ctx, ObservationSide::Write)
506            .is_ready()
507        {
508            ctx.waker().wake_by_ref();
509        }
510        Poll::Pending
511    }
512}
513
514impl<S: ExtensionsRef> ExtensionsRef for HARWebSocket<S> {
515    fn extensions(&self) -> &Extensions {
516        self.inner.extensions()
517    }
518}
519
520impl<S> Stream for HARWebSocket<S>
521where
522    S: Stream<Item = Result<Message, ProtocolError>> + Unpin,
523{
524    type Item = Result<Message, ProtocolError>;
525
526    fn poll_next(self: Pin<&mut Self>, ctx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
527        let this = self.get_mut();
528        ready!(this.poll_observation(ctx, ObservationSide::Read));
529        if let Some(message) = this.pending_read.take() {
530            return Poll::Ready(Some(message));
531        }
532
533        match ready!(Pin::new(&mut this.inner).poll_next(ctx)) {
534            Some(Ok(message)) => {
535                let close_after = matches!(&message, Message::Close(_));
536                if this.begin_message_observation(false, &message, close_after) {
537                    this.pending_read = Some(Ok(message));
538                    ready!(this.poll_observation(ctx, ObservationSide::Read));
539                    Poll::Ready(this.pending_read.take())
540                } else {
541                    Poll::Ready(Some(Ok(message)))
542                }
543            }
544            Some(Err(error)) => {
545                this.begin_error_observation(&error);
546                if this.pending_observation.is_some() {
547                    this.pending_read = Some(Err(error));
548                    ready!(this.poll_observation(ctx, ObservationSide::Read));
549                    Poll::Ready(this.pending_read.take())
550                } else {
551                    Poll::Ready(Some(Err(error)))
552                }
553            }
554            None => {
555                if this.close_on_terminal {
556                    if let Some(capture) = &this.capture_lease {
557                        capture.close();
558                    }
559                    this.capture_lease.take();
560                }
561                Poll::Ready(None)
562            }
563        }
564    }
565}
566
567impl<S> Sink<Message> for HARWebSocket<S>
568where
569    S: Sink<Message, Error = ProtocolError> + Unpin,
570{
571    type Error = ProtocolError;
572
573    fn poll_ready(self: Pin<&mut Self>, ctx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
574        let this = self.get_mut();
575        ready!(this.poll_observation(ctx, ObservationSide::Write));
576        if let Some(error) = this.pending_write_error.take() {
577            return Poll::Ready(Err(error));
578        }
579        match Pin::new(&mut this.inner).poll_ready(ctx) {
580            Poll::Ready(Err(error)) => this.poll_write_error(ctx, error),
581            result => result,
582        }
583    }
584
585    fn start_send(self: Pin<&mut Self>, item: Message) -> Result<(), Self::Error> {
586        let this = self.get_mut();
587        let observation = this.message_observation(true, &item);
588        match Pin::new(&mut this.inner).start_send(item) {
589            Ok(()) => {
590                if let Some(observation) = observation {
591                    this.set_pending_observation(observation);
592                }
593                Ok(())
594            }
595            Err(error) => {
596                drop(observation);
597                if this.begin_error_observation(&error) {
598                    this.pending_write_error = Some(error);
599                    Ok(())
600                } else {
601                    Err(error)
602                }
603            }
604        }
605    }
606
607    fn poll_flush(self: Pin<&mut Self>, ctx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
608        let this = self.get_mut();
609        ready!(this.poll_observation(ctx, ObservationSide::Write));
610        if let Some(error) = this.pending_write_error.take() {
611            return Poll::Ready(Err(error));
612        }
613        match Pin::new(&mut this.inner).poll_flush(ctx) {
614            Poll::Ready(Err(error)) => this.poll_write_error(ctx, error),
615            result => result,
616        }
617    }
618
619    fn poll_close(self: Pin<&mut Self>, ctx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
620        let this = self.get_mut();
621        ready!(this.poll_observation(ctx, ObservationSide::Write));
622        if let Some(error) = this.pending_write_error.take() {
623            return Poll::Ready(Err(error));
624        }
625        match Pin::new(&mut this.inner).poll_close(ctx) {
626            Poll::Ready(Err(error)) => this.poll_write_error(ctx, error),
627            result => result,
628        }
629    }
630}
631
632impl<S> HARWebSocket<S>
633where
634    S: Stream<Item = Result<Message, ProtocolError>> + Sink<Message, Error = ProtocolError> + Unpin,
635{
636    /// Write and flush one message.
637    pub async fn send_message(&mut self, message: Message) -> Result<(), ProtocolError> {
638        self.send(message).await
639    }
640
641    /// Receive one complete message.
642    pub async fn recv_message(&mut self) -> Result<Message, ProtocolError> {
643        self.next().await.ok_or_else(|| {
644            ProtocolError::Io(io::Error::new(
645                io::ErrorKind::ConnectionAborted,
646                "Connection closed: no messages to receive",
647            ))
648        })?
649    }
650
651    /// Close the WebSocket.
652    pub async fn close(
653        &mut self,
654        message: Option<crate::protocol::CloseFrame>,
655    ) -> Result<(), ProtocolError> {
656        self.send(Message::Close(message)).await
657    }
658}
659
660fn into_har_message(
661    message_type: WebSocketMessageType,
662    message: &Message,
663) -> Option<WebSocketMessage> {
664    let time = epoch_seconds_from_millis(unix_timestamp_millis());
665    match message {
666        Message::Text(data) => Some(WebSocketMessage::text(message_type, time, data.as_str())),
667        Message::Binary(data) => Some(WebSocketMessage::binary(message_type, time, data)),
668        Message::Ping(_) | Message::Pong(_) | Message::Close(_) | Message::Frame(_) => None,
669    }
670}
671
672fn epoch_seconds_from_millis(timestamp: i64) -> f64 {
673    timestamp as f64 / 1_000.0
674}
675
676#[cfg(test)]
677mod tests {
678    use super::{
679        HARWebSocket, HARWebSocketLayer, ObservationSide, epoch_seconds_from_millis,
680        into_har_message,
681    };
682    use crate::{
683        AsyncWebSocket, Message,
684        handshake::mitm::WebSocketBridge,
685        protocol::{Role, WebSocketConfig, frame::Frame},
686    };
687    use parking_lot::Mutex;
688    use rama_core::{
689        Layer, Service, ServiceInput,
690        error::BoxError,
691        extensions::{Extensions, ExtensionsRef},
692        futures::{Sink, SinkExt as _, Stream, StreamExt as _},
693        service::service_fn,
694    };
695    use rama_http::layer::har::{
696        recorder::{WebSocketCapture, WebSocketCaptureRecorder},
697        spec::{WebSocketMessage, WebSocketMessageOpcode, WebSocketMessageType},
698    };
699    use std::{
700        future::Future,
701        io,
702        pin::Pin,
703        sync::{
704            Arc,
705            atomic::{AtomicBool, AtomicUsize, Ordering},
706        },
707        task::{Context, Poll, Wake, Waker},
708    };
709    use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
710    use tokio::sync::Notify;
711
712    #[derive(Default)]
713    struct TestState {
714        messages: Mutex<Vec<WebSocketMessage>>,
715        closes: AtomicUsize,
716    }
717
718    struct TestRecorder(Arc<TestState>);
719
720    impl WebSocketCaptureRecorder for TestRecorder {
721        async fn record(&self, message: WebSocketMessage) -> Result<(), BoxError> {
722            self.0.messages.lock().push(message);
723            Ok(())
724        }
725    }
726
727    #[derive(Debug, Clone, Copy, PartialEq, Eq, rama_core::extensions::Extension)]
728    struct TestExtension(u8);
729
730    #[derive(Debug, Default)]
731    struct DelegatingSocketState {
732        ready: AtomicUsize,
733        closes: AtomicUsize,
734        messages: Mutex<Vec<Message>>,
735    }
736
737    #[derive(Debug)]
738    struct DelegatingSocket {
739        extensions: Extensions,
740        state: Arc<DelegatingSocketState>,
741    }
742
743    impl ExtensionsRef for DelegatingSocket {
744        fn extensions(&self) -> &Extensions {
745            &self.extensions
746        }
747    }
748
749    impl Stream for DelegatingSocket {
750        type Item = Result<Message, crate::ProtocolError>;
751
752        fn poll_next(self: Pin<&mut Self>, _ctx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
753            Poll::Pending
754        }
755    }
756
757    impl Sink<Message> for DelegatingSocket {
758        type Error = crate::ProtocolError;
759
760        fn poll_ready(
761            self: Pin<&mut Self>,
762            _ctx: &mut Context<'_>,
763        ) -> Poll<Result<(), Self::Error>> {
764            self.state.ready.fetch_add(1, Ordering::AcqRel);
765            Poll::Ready(Ok(()))
766        }
767
768        fn start_send(self: Pin<&mut Self>, item: Message) -> Result<(), Self::Error> {
769            self.state.messages.lock().push(item);
770            Ok(())
771        }
772
773        fn poll_flush(
774            self: Pin<&mut Self>,
775            _ctx: &mut Context<'_>,
776        ) -> Poll<Result<(), Self::Error>> {
777            Poll::Ready(Ok(()))
778        }
779
780        fn poll_close(
781            self: Pin<&mut Self>,
782            _ctx: &mut Context<'_>,
783        ) -> Poll<Result<(), Self::Error>> {
784            self.state.closes.fetch_add(1, Ordering::AcqRel);
785            Poll::Ready(Ok(()))
786        }
787    }
788
789    struct TailSocket {
790        extensions: Extensions,
791        incoming: Option<Message>,
792        sent: Arc<Mutex<Vec<Message>>>,
793    }
794
795    impl ExtensionsRef for TailSocket {
796        fn extensions(&self) -> &Extensions {
797            &self.extensions
798        }
799    }
800
801    impl Stream for TailSocket {
802        type Item = Result<Message, crate::ProtocolError>;
803
804        fn poll_next(mut self: Pin<&mut Self>, _ctx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
805            Poll::Ready(self.incoming.take().map(Ok))
806        }
807    }
808
809    impl Sink<Message> for TailSocket {
810        type Error = crate::ProtocolError;
811
812        fn poll_ready(
813            self: Pin<&mut Self>,
814            _ctx: &mut Context<'_>,
815        ) -> Poll<Result<(), Self::Error>> {
816            Poll::Ready(Ok(()))
817        }
818
819        fn start_send(self: Pin<&mut Self>, item: Message) -> Result<(), Self::Error> {
820            self.sent.lock().push(item);
821            Ok(())
822        }
823
824        fn poll_flush(
825            self: Pin<&mut Self>,
826            _ctx: &mut Context<'_>,
827        ) -> Poll<Result<(), Self::Error>> {
828            Poll::Ready(Ok(()))
829        }
830
831        fn poll_close(
832            self: Pin<&mut Self>,
833            _ctx: &mut Context<'_>,
834        ) -> Poll<Result<(), Self::Error>> {
835            Poll::Ready(Ok(()))
836        }
837    }
838
839    #[derive(Clone, Copy)]
840    enum SinkFailurePoint {
841        Ready,
842        Flush,
843        Close,
844    }
845
846    struct FailingSink(SinkFailurePoint);
847
848    impl Stream for FailingSink {
849        type Item = Result<Message, crate::ProtocolError>;
850
851        fn poll_next(self: Pin<&mut Self>, _ctx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
852            Poll::Pending
853        }
854    }
855
856    impl FailingSink {
857        fn error(&self) -> crate::ProtocolError {
858            crate::ProtocolError::Io(io::Error::other(match self.0 {
859                SinkFailurePoint::Ready => "ready failed",
860                SinkFailurePoint::Flush => "flush failed",
861                SinkFailurePoint::Close => "close failed",
862            }))
863        }
864    }
865
866    impl Sink<Message> for FailingSink {
867        type Error = crate::ProtocolError;
868
869        fn poll_ready(
870            self: Pin<&mut Self>,
871            _ctx: &mut Context<'_>,
872        ) -> Poll<Result<(), Self::Error>> {
873            match self.0 {
874                SinkFailurePoint::Ready => Poll::Ready(Err(self.error())),
875                _ => Poll::Ready(Ok(())),
876            }
877        }
878
879        fn start_send(self: Pin<&mut Self>, _item: Message) -> Result<(), Self::Error> {
880            Ok(())
881        }
882
883        fn poll_flush(
884            self: Pin<&mut Self>,
885            _ctx: &mut Context<'_>,
886        ) -> Poll<Result<(), Self::Error>> {
887            match self.0 {
888                SinkFailurePoint::Flush => Poll::Ready(Err(self.error())),
889                _ => Poll::Ready(Ok(())),
890            }
891        }
892
893        fn poll_close(
894            self: Pin<&mut Self>,
895            _ctx: &mut Context<'_>,
896        ) -> Poll<Result<(), Self::Error>> {
897            match self.0 {
898                SinkFailurePoint::Close => Poll::Ready(Err(self.error())),
899                _ => Poll::Ready(Ok(())),
900            }
901        }
902    }
903
904    #[derive(Default)]
905    struct ReadinessState {
906        ready: AtomicBool,
907        polls: AtomicUsize,
908        notify: Notify,
909        messages: Mutex<Vec<WebSocketMessage>>,
910    }
911
912    struct StallingRecorder(Arc<ReadinessState>);
913
914    impl WebSocketCaptureRecorder for StallingRecorder {
915        async fn record(&self, message: WebSocketMessage) -> Result<(), BoxError> {
916            loop {
917                let notified = self.0.notify.notified();
918                if self.0.ready.swap(false, Ordering::AcqRel) {
919                    break;
920                }
921                tokio::pin!(notified);
922                std::future::poll_fn(|ctx| {
923                    self.0.polls.fetch_add(1, Ordering::AcqRel);
924                    notified.as_mut().poll(ctx)
925                })
926                .await;
927            }
928            self.0.messages.lock().push(message);
929            Ok(())
930        }
931    }
932
933    struct FailingRecorder(Arc<AtomicUsize>);
934
935    impl WebSocketCaptureRecorder for FailingRecorder {
936        async fn record(&self, _message: WebSocketMessage) -> Result<(), BoxError> {
937            self.0.fetch_add(1, Ordering::AcqRel);
938            Err(io::Error::other("recorder failed").into())
939        }
940    }
941
942    #[derive(Default)]
943    struct WakeCounter(AtomicUsize);
944
945    impl Wake for WakeCounter {
946        fn wake(self: Arc<Self>) {
947            self.0.fetch_add(1, Ordering::AcqRel);
948        }
949
950        fn wake_by_ref(self: &Arc<Self>) {
951            self.0.fetch_add(1, Ordering::AcqRel);
952        }
953    }
954
955    #[test]
956    fn observation_waker_wakes_both_sides_by_ref() {
957        let observation_wakers = Arc::new(super::ObservationWakers::new());
958        let read = Arc::new(WakeCounter::default());
959        let write = Arc::new(WakeCounter::default());
960        observation_wakers.register(ObservationSide::Read, &Waker::from(read.clone()));
961        observation_wakers.register(ObservationSide::Write, &Waker::from(write.clone()));
962
963        Waker::from(observation_wakers).wake_by_ref();
964
965        assert_eq!(read.0.load(Ordering::Acquire), 1);
966        assert_eq!(write.0.load(Ordering::Acquire), 1);
967    }
968
969    #[derive(Clone, Copy)]
970    enum WriteBehavior {
971        Pending,
972        BrokenPipe,
973    }
974
975    struct TestIo(WriteBehavior);
976
977    impl AsyncRead for TestIo {
978        fn poll_read(
979            self: Pin<&mut Self>,
980            _ctx: &mut Context<'_>,
981            _buf: &mut ReadBuf<'_>,
982        ) -> Poll<io::Result<()>> {
983            Poll::Pending
984        }
985    }
986
987    impl AsyncWrite for TestIo {
988        fn poll_write(
989            self: Pin<&mut Self>,
990            _ctx: &mut Context<'_>,
991            _buf: &[u8],
992        ) -> Poll<io::Result<usize>> {
993            match self.0 {
994                WriteBehavior::Pending => Poll::Pending,
995                WriteBehavior::BrokenPipe => {
996                    Poll::Ready(Err(io::Error::from(io::ErrorKind::BrokenPipe)))
997                }
998            }
999        }
1000
1001        fn poll_flush(self: Pin<&mut Self>, _ctx: &mut Context<'_>) -> Poll<io::Result<()>> {
1002            Poll::Ready(Ok(()))
1003        }
1004
1005        fn poll_shutdown(self: Pin<&mut Self>, _ctx: &mut Context<'_>) -> Poll<io::Result<()>> {
1006            Poll::Ready(Ok(()))
1007        }
1008    }
1009
1010    async fn socket_with_write_behavior(
1011        behavior: WriteBehavior,
1012        state: Arc<TestState>,
1013    ) -> HARWebSocket<AsyncWebSocket<ServiceInput<TestIo>>> {
1014        let socket = AsyncWebSocket::from_raw_socket(
1015            ServiceInput::new(TestIo(behavior)),
1016            Role::Client,
1017            Some(WebSocketConfig::default().with_write_buffer_size(0)),
1018        )
1019        .await;
1020        HARWebSocket::new(
1021            socket,
1022            Role::Client,
1023            Some(WebSocketCapture::new(
1024                TestRecorder(state.clone()),
1025                move || {
1026                    state.closes.fetch_add(1, Ordering::AcqRel);
1027                },
1028            )),
1029        )
1030    }
1031
1032    #[tokio::test]
1033    async fn start_send_distinguishes_backpressure_from_fatal_io() {
1034        let pending_sink = Arc::new(TestState::default());
1035        let mut pending =
1036            socket_with_write_behavior(WriteBehavior::Pending, pending_sink.clone()).await;
1037        std::future::poll_fn(|ctx| Sink::poll_ready(Pin::new(&mut pending), ctx))
1038            .await
1039            .expect("pending socket ready");
1040        Sink::start_send(Pin::new(&mut pending), Message::text("queued"))
1041            .expect("WouldBlock means the frame was accepted into the write buffer");
1042        std::future::poll_fn(|ctx| pending.poll_observation(ctx, ObservationSide::Write)).await;
1043        {
1044            let pending_messages = pending_sink.messages.lock();
1045            assert_eq!(pending_messages.len(), 1);
1046            assert_eq!(pending_messages[0].r#type, WebSocketMessageType::Send);
1047            assert_eq!(pending_messages[0].data.as_str(), "queued");
1048        }
1049
1050        let broken_sink = Arc::new(TestState::default());
1051        let mut broken =
1052            socket_with_write_behavior(WriteBehavior::BrokenPipe, broken_sink.clone()).await;
1053        broken
1054            .send_message(Message::text("rejected"))
1055            .await
1056            .expect_err("normal send flow returns the transport error after recording it");
1057        let broken_messages = broken_sink.messages.lock();
1058        assert_eq!(broken_messages.len(), 1);
1059        assert_eq!(broken_messages[0].r#type, WebSocketMessageType::Error);
1060        assert_eq!(broken_messages[0].opcode, WebSocketMessageOpcode::ERROR);
1061        drop(broken_messages);
1062        assert_eq!(broken_sink.closes.load(Ordering::Acquire), 1);
1063    }
1064
1065    #[tokio::test]
1066    async fn sink_poll_errors_are_recorded_before_being_returned() {
1067        for failure in [
1068            SinkFailurePoint::Ready,
1069            SinkFailurePoint::Flush,
1070            SinkFailurePoint::Close,
1071        ] {
1072            let state = Arc::new(TestState::default());
1073            let mut socket = HARWebSocket::new(
1074                FailingSink(failure),
1075                Role::Client,
1076                Some(WebSocketCapture::new(TestRecorder(state.clone()), {
1077                    let state = state.clone();
1078                    move || {
1079                        state.closes.fetch_add(1, Ordering::AcqRel);
1080                    }
1081                })),
1082            );
1083
1084            let result = match failure {
1085                SinkFailurePoint::Ready => {
1086                    std::future::poll_fn(|ctx| Sink::poll_ready(Pin::new(&mut socket), ctx)).await
1087                }
1088                SinkFailurePoint::Flush => {
1089                    std::future::poll_fn(|ctx| Sink::poll_flush(Pin::new(&mut socket), ctx)).await
1090                }
1091                SinkFailurePoint::Close => {
1092                    std::future::poll_fn(|ctx| Sink::poll_close(Pin::new(&mut socket), ctx)).await
1093                }
1094            };
1095
1096            assert!(result.is_err());
1097            let messages = state.messages.lock();
1098            assert_eq!(messages.len(), 1);
1099            assert_eq!(messages[0].r#type, WebSocketMessageType::Error);
1100            assert_eq!(messages[0].opcode, WebSocketMessageOpcode::ERROR);
1101            drop(messages);
1102            assert_eq!(state.closes.load(Ordering::Acquire), 1);
1103        }
1104    }
1105
1106    #[tokio::test]
1107    async fn closing_sink_keeps_capture_alive_for_tail_reads() {
1108        let state = Arc::new(TestState::default());
1109        let mut socket = HARWebSocket::new(
1110            TailSocket {
1111                extensions: Extensions::new(),
1112                incoming: Some(Message::text("tail")),
1113                sent: Arc::new(Mutex::new(Vec::new())),
1114            },
1115            Role::Client,
1116            Some(WebSocketCapture::new(TestRecorder(state.clone()), {
1117                let state = state.clone();
1118                move || {
1119                    state.closes.fetch_add(1, Ordering::AcqRel);
1120                }
1121            })),
1122        );
1123
1124        std::future::poll_fn(|ctx| Sink::poll_close(Pin::new(&mut socket), ctx))
1125            .await
1126            .expect("close write half");
1127        assert_eq!(state.closes.load(Ordering::Acquire), 0);
1128        match socket.next().await {
1129            Some(Ok(message)) => assert_eq!(message, Message::text("tail")),
1130            other => panic!("unexpected tail read: {other:?}"),
1131        }
1132        assert!(socket.next().await.is_none());
1133
1134        let messages = state.messages.lock();
1135        assert_eq!(messages.len(), 1);
1136        assert_eq!(messages[0].r#type, WebSocketMessageType::Receive);
1137        assert_eq!(messages[0].data.as_str(), "tail");
1138        drop(messages);
1139        assert_eq!(state.closes.load(Ordering::Acquire), 1);
1140    }
1141
1142    #[tokio::test]
1143    async fn legal_stream_sink_interleave_preserves_both_observations() {
1144        let state = Arc::new(ReadinessState::default());
1145        let mut socket = HARWebSocket::new(
1146            TailSocket {
1147                extensions: Extensions::new(),
1148                incoming: Some(Message::text("incoming")),
1149                sent: Arc::new(Mutex::new(Vec::new())),
1150            },
1151            Role::Client,
1152            Some(WebSocketCapture::new(
1153                StallingRecorder(state.clone()),
1154                || {},
1155            )),
1156        );
1157
1158        let waker = Waker::noop();
1159        let mut ctx = Context::from_waker(waker);
1160        assert!(Sink::poll_ready(Pin::new(&mut socket), &mut ctx).is_ready());
1161        assert!(Stream::poll_next(Pin::new(&mut socket), &mut ctx).is_pending());
1162        Sink::start_send(Pin::new(&mut socket), Message::text("outgoing"))
1163            .expect("send after earlier readiness");
1164
1165        state.ready.store(true, Ordering::Release);
1166        state.notify.notify_one();
1167        assert!(Sink::poll_ready(Pin::new(&mut socket), &mut ctx).is_pending());
1168        assert_eq!(state.messages.lock().len(), 1);
1169
1170        state.ready.store(true, Ordering::Release);
1171        state.notify.notify_one();
1172        std::future::poll_fn(|ctx| Sink::poll_ready(Pin::new(&mut socket), ctx))
1173            .await
1174            .expect("both observations finish before readiness");
1175        match Stream::poll_next(Pin::new(&mut socket), &mut ctx) {
1176            Poll::Ready(Some(Ok(message))) => {
1177                assert_eq!(message, Message::text("incoming"));
1178            }
1179            other => panic!("unexpected pending read: {other:?}"),
1180        }
1181
1182        let messages = state.messages.lock();
1183        assert_eq!(messages.len(), 2);
1184        assert_eq!(messages[0].r#type, WebSocketMessageType::Receive);
1185        assert_eq!(messages[0].data.as_str(), "incoming");
1186        assert_eq!(messages[1].r#type, WebSocketMessageType::Send);
1187        assert_eq!(messages[1].data.as_str(), "outgoing");
1188    }
1189
1190    #[tokio::test]
1191    async fn recorder_failure_detaches_capture_without_failing_socket() {
1192        let attempts = Arc::new(AtomicUsize::new(0));
1193        let closes = Arc::new(AtomicUsize::new(0));
1194        let state = Arc::new(DelegatingSocketState::default());
1195        let mut socket = HARWebSocket::new(
1196            DelegatingSocket {
1197                extensions: Extensions::new(),
1198                state: state.clone(),
1199            },
1200            Role::Client,
1201            Some(WebSocketCapture::new(FailingRecorder(attempts.clone()), {
1202                let closes = closes.clone();
1203                move || {
1204                    closes.fetch_add(1, Ordering::AcqRel);
1205                }
1206            })),
1207        );
1208
1209        socket
1210            .send_message(Message::text("still-forwarded"))
1211            .await
1212            .expect("capture failure does not fail the WebSocket");
1213        socket
1214            .send_message(Message::text("capture-detached"))
1215            .await
1216            .expect("subsequent traffic bypasses failed capture");
1217
1218        assert_eq!(attempts.load(Ordering::Acquire), 1);
1219        assert_eq!(closes.load(Ordering::Acquire), 1);
1220        assert!(socket.capture_lease.is_none());
1221        assert_eq!(state.messages.lock().len(), 2);
1222    }
1223
1224    #[tokio::test]
1225    async fn async_recorder_backpressures_web_socket_sends() {
1226        let sink = Arc::new(ReadinessState::default());
1227        let socket = AsyncWebSocket::from_raw_socket(
1228            ServiceInput::new(TestIo(WriteBehavior::Pending)),
1229            Role::Client,
1230            None,
1231        )
1232        .await;
1233        let mut socket = HARWebSocket::new(
1234            socket,
1235            Role::Client,
1236            Some(WebSocketCapture::new(StallingRecorder(sink.clone()), || {})),
1237        );
1238        std::future::poll_fn(|ctx| Sink::poll_ready(Pin::new(&mut socket), ctx))
1239            .await
1240            .expect("socket initially ready");
1241        Sink::start_send(Pin::new(&mut socket), Message::text("bounded"))
1242            .expect("socket accepts message before recording it");
1243
1244        let mut observation = Box::pin(std::future::poll_fn(|ctx| {
1245            socket.poll_observation(ctx, ObservationSide::Write)
1246        }));
1247        assert!(rama_core::futures::poll!(&mut observation).is_pending());
1248        sink.ready.store(true, Ordering::Release);
1249        sink.notify.notify_one();
1250        observation.await;
1251
1252        let messages = sink.messages.lock();
1253        assert_eq!(messages.len(), 1);
1254        assert_eq!(messages[0].data.as_str(), "bounded");
1255    }
1256
1257    #[tokio::test]
1258    async fn async_recorder_backpressures_incoming_web_socket_messages() {
1259        let sink = Arc::new(ReadinessState::default());
1260        let (server_io, client_io) = tokio::io::duplex(1024);
1261        let server =
1262            AsyncWebSocket::from_raw_socket(ServiceInput::new(server_io), Role::Server, None).await;
1263        let mut server = HARWebSocket::new(
1264            server,
1265            Role::Server,
1266            Some(WebSocketCapture::new(StallingRecorder(sink.clone()), || {})),
1267        );
1268        let mut client =
1269            AsyncWebSocket::from_raw_socket(ServiceInput::new(client_io), Role::Client, None).await;
1270
1271        client
1272            .send_message(Message::text("incoming"))
1273            .await
1274            .expect("send test message");
1275        let mut receive = Box::pin(server.recv_message());
1276        assert!(rama_core::futures::poll!(&mut receive).is_pending());
1277        sink.ready.store(true, Ordering::Release);
1278        sink.notify.notify_one();
1279
1280        assert_eq!(receive.await.unwrap(), Message::text("incoming"));
1281        let messages = sink.messages.lock();
1282        assert_eq!(messages.len(), 1);
1283        assert_eq!(messages[0].data.as_str(), "incoming");
1284    }
1285
1286    #[tokio::test]
1287    async fn split_socket_keeps_independent_recorder_wakers() {
1288        let recorder_state = Arc::new(ReadinessState::default());
1289        let socket_state = Arc::new(DelegatingSocketState::default());
1290        let socket = HARWebSocket::new(
1291            DelegatingSocket {
1292                extensions: Extensions::new(),
1293                state: socket_state.clone(),
1294            },
1295            Role::Client,
1296            Some(WebSocketCapture::new(
1297                StallingRecorder(recorder_state.clone()),
1298                || {},
1299            )),
1300        );
1301        let (mut writer, mut reader) = socket.split();
1302
1303        let writer_task = tokio::spawn(async move { writer.send(Message::text("split")).await });
1304        tokio::time::timeout(std::time::Duration::from_secs(1), async {
1305            while recorder_state.polls.load(Ordering::Acquire) == 0 {
1306                tokio::task::yield_now().await;
1307            }
1308        })
1309        .await
1310        .expect("writer polls the recorder");
1311
1312        let reader_task = tokio::spawn(async move { reader.next().await });
1313        tokio::time::timeout(std::time::Duration::from_secs(1), async {
1314            while recorder_state.polls.load(Ordering::Acquire) < 2 {
1315                tokio::task::yield_now().await;
1316            }
1317        })
1318        .await
1319        .expect("reader repolls the pending recorder future");
1320
1321        // The reader was the last task to poll the recorder future. Cancelling
1322        // it must not strand the writer's earlier waiter.
1323        reader_task.abort();
1324        _ = reader_task.await;
1325
1326        recorder_state.ready.store(true, Ordering::Release);
1327        recorder_state.notify.notify_one();
1328        tokio::time::timeout(std::time::Duration::from_secs(1), writer_task)
1329            .await
1330            .expect("split writer is woken after recorder completion")
1331            .expect("writer task succeeds")
1332            .expect("split send succeeds");
1333    }
1334
1335    #[tokio::test]
1336    async fn relay_layer_claims_only_egress_capture() {
1337        let ingress_capture =
1338            WebSocketCapture::new(TestRecorder(Arc::new(TestState::default())), || {});
1339        let egress_capture =
1340            WebSocketCapture::new(TestRecorder(Arc::new(TestState::default())), || {});
1341
1342        let ingress_extensions = Extensions::new();
1343        ingress_extensions.insert(ingress_capture.clone());
1344        let egress_extensions = Extensions::new();
1345        egress_extensions.insert(egress_capture.clone());
1346
1347        let inner = service_fn(
1348            |bridge: WebSocketBridge<
1349                HARWebSocket<DelegatingSocket>,
1350                HARWebSocket<DelegatingSocket>,
1351            >| async move {
1352                assert!(bridge.ingress.capture_lease.is_some());
1353                assert!(bridge.egress.capture_lease.is_some());
1354                Ok::<_, std::convert::Infallible>(())
1355            },
1356        );
1357        HARWebSocketLayer::new()
1358            .into_layer(inner)
1359            .serve(WebSocketBridge {
1360                ingress: DelegatingSocket {
1361                    extensions: ingress_extensions,
1362                    state: Arc::new(DelegatingSocketState::default()),
1363                },
1364                egress: DelegatingSocket {
1365                    extensions: egress_extensions,
1366                    state: Arc::new(DelegatingSocketState::default()),
1367                },
1368            })
1369            .await
1370            .expect("HAR relay layer is infallible");
1371
1372        let ingress_lease = ingress_capture
1373            .lease()
1374            .expect("ingress capture remains unclaimed");
1375        assert!(
1376            egress_capture.lease().is_none(),
1377            "egress capture was claimed for the relay"
1378        );
1379        drop(ingress_lease);
1380    }
1381
1382    #[tokio::test]
1383    async fn relay_capture_lives_as_long_as_returned_bridge() {
1384        let state = Arc::new(TestState::default());
1385        let capture = WebSocketCapture::new(TestRecorder(state.clone()), {
1386            let state = state.clone();
1387            move || {
1388                state.closes.fetch_add(1, Ordering::AcqRel);
1389            }
1390        });
1391        let egress_extensions = Extensions::new();
1392        egress_extensions.insert(capture.clone());
1393
1394        let mut bridge = HARWebSocketLayer::new()
1395            .into_layer(())
1396            .serve(WebSocketBridge {
1397                ingress: DelegatingSocket {
1398                    extensions: Extensions::new(),
1399                    state: Arc::new(DelegatingSocketState::default()),
1400                },
1401                egress: DelegatingSocket {
1402                    extensions: egress_extensions,
1403                    state: Arc::new(DelegatingSocketState::default()),
1404                },
1405            })
1406            .await
1407            .expect("identity service returns the decorated bridge");
1408
1409        assert_eq!(state.closes.load(Ordering::Acquire), 0);
1410        bridge
1411            .egress
1412            .send_message(Message::text("after-service-return"))
1413            .await
1414            .expect("live bridge keeps recording");
1415        assert_eq!(state.messages.lock().len(), 1);
1416        drop(bridge);
1417        assert_eq!(state.closes.load(Ordering::Acquire), 1);
1418    }
1419
1420    #[test]
1421    fn explicitly_closed_capture_skips_message_conversion() {
1422        let capture = WebSocketCapture::new(TestRecorder(Arc::new(TestState::default())), || {});
1423        let mut socket = HARWebSocket::new(
1424            DelegatingSocket {
1425                extensions: Extensions::new(),
1426                state: Arc::new(DelegatingSocketState::default()),
1427            },
1428            Role::Client,
1429            Some(capture.clone()),
1430        );
1431
1432        capture.close();
1433        assert!(
1434            socket
1435                .message_observation(true, &Message::binary(vec![0; 1024]))
1436                .is_none()
1437        );
1438        assert!(
1439            socket
1440                .capture_lease
1441                .as_ref()
1442                .is_some_and(|lease| lease.is_closed())
1443        );
1444        socket.capture_lease.take();
1445    }
1446
1447    #[tokio::test]
1448    async fn server_role_uses_client_perspective() {
1449        let sink = Arc::new(TestState::default());
1450        let socket = AsyncWebSocket::from_raw_socket(
1451            ServiceInput::new(tokio::io::duplex(1024).0),
1452            Role::Server,
1453            None,
1454        )
1455        .await;
1456        let mut socket = HARWebSocket::new(
1457            socket,
1458            Role::Server,
1459            Some(WebSocketCapture::new(TestRecorder(sink.clone()), || {})),
1460        );
1461
1462        assert!(socket.begin_message_observation(false, &Message::text("from-client"), false));
1463        std::future::poll_fn(|ctx| socket.poll_observation(ctx, ObservationSide::Read)).await;
1464        assert!(socket.begin_message_observation(true, &Message::binary(vec![1, 2]), false));
1465        std::future::poll_fn(|ctx| socket.poll_observation(ctx, ObservationSide::Write)).await;
1466
1467        let messages = sink.messages.lock();
1468        assert_eq!(messages.len(), 2);
1469        assert_eq!(messages[0].r#type, WebSocketMessageType::Send);
1470        assert_eq!(messages[0].opcode, WebSocketMessageOpcode::TEXT);
1471        assert_eq!(messages[1].r#type, WebSocketMessageType::Receive);
1472        assert_eq!(messages[1].opcode, WebSocketMessageOpcode::BINARY);
1473    }
1474
1475    #[tokio::test]
1476    async fn wrapper_delegates_socket_contract_and_convenience_methods() {
1477        let extensions = Extensions::new();
1478        extensions.insert(TestExtension(42));
1479        let state = Arc::new(DelegatingSocketState::default());
1480        let mut socket = HARWebSocket::new(
1481            DelegatingSocket {
1482                extensions,
1483                state: state.clone(),
1484            },
1485            Role::Client,
1486            None,
1487        );
1488
1489        assert_eq!(
1490            socket.extensions().get_ref::<TestExtension>(),
1491            Some(&TestExtension(42))
1492        );
1493        std::future::poll_fn(|ctx| Sink::poll_ready(Pin::new(&mut socket), ctx))
1494            .await
1495            .expect("inner sink ready");
1496        socket
1497            .send_message(Message::text("message"))
1498            .await
1499            .expect("send convenience method delegates");
1500        socket
1501            .close(None)
1502            .await
1503            .expect("close convenience method delegates");
1504        std::future::poll_fn(|ctx| Sink::poll_close(Pin::new(&mut socket), ctx))
1505            .await
1506            .expect("inner sink closes");
1507
1508        assert_eq!(state.ready.load(Ordering::Acquire), 3);
1509        assert_eq!(state.closes.load(Ordering::Acquire), 1);
1510        assert_eq!(
1511            *state.messages.lock(),
1512            vec![Message::text("message"), Message::Close(None)]
1513        );
1514        assert!(format!("{socket:?}").contains("HARWebSocket"));
1515    }
1516
1517    #[test]
1518    fn pending_observation_debug_exposes_capture_state() {
1519        let capture = WebSocketCapture::new(TestRecorder(Arc::new(TestState::default())), || {});
1520        let lease = capture.lease().expect("capture lease");
1521        let observation = super::PendingObservation {
1522            future: lease.record(WebSocketMessage::text(
1523                WebSocketMessageType::Send,
1524                1.0,
1525                "message",
1526            )),
1527            close_after: true,
1528        };
1529
1530        let debug = format!("{observation:?}");
1531        assert!(debug.contains("PendingObservation"));
1532        assert!(debug.contains("close_after: true"));
1533    }
1534
1535    #[test]
1536    fn har_messages_encode_complete_data_messages() {
1537        let cases = [
1538            (
1539                Message::text("hello"),
1540                WebSocketMessageOpcode::TEXT,
1541                "hello",
1542            ),
1543            (
1544                Message::binary(vec![0_u8, 1, 0xff]),
1545                WebSocketMessageOpcode::BINARY,
1546                "AAH/",
1547            ),
1548        ];
1549
1550        for (message, opcode, data) in cases {
1551            let message = into_har_message(WebSocketMessageType::Send, &message)
1552                .expect("complete data message");
1553            assert_eq!(message.r#type, WebSocketMessageType::Send);
1554            assert_eq!(message.opcode, opcode);
1555            assert_eq!(message.data.as_str(), data);
1556            assert!(message.time > 1_700_000_000.0);
1557        }
1558    }
1559
1560    #[test]
1561    fn har_messages_skip_control_and_raw_frames() {
1562        for message in [
1563            Message::Ping(vec![2, 3].into()),
1564            Message::Pong(vec![4, 5].into()),
1565            Message::Close(None),
1566            Message::Frame(Frame::ping(rama_core::bytes::Bytes::from_static(&[6]))),
1567        ] {
1568            assert!(into_har_message(WebSocketMessageType::Send, &message).is_none());
1569        }
1570    }
1571
1572    #[test]
1573    fn har_timestamp_conversion_preserves_milliseconds() {
1574        assert_eq!(
1575            epoch_seconds_from_millis(1_558_730_482_507),
1576            1_558_730_482.507
1577        );
1578    }
1579}