1use 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
84struct 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#[derive(Debug, Clone, Copy, Default)]
126pub struct HARWebSocketLayer;
127
128impl HARWebSocketLayer {
129 #[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#[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 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
266pub 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 #[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 #[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 #[must_use]
363 pub fn into_inner(self) -> S {
364 self.inner
365 }
366
367 #[must_use]
369 pub fn get_ref(&self) -> &S {
370 &self.inner
371 }
372
373 #[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 pub async fn send_message(&mut self, message: Message) -> Result<(), ProtocolError> {
638 self.send(message).await
639 }
640
641 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 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 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}