1use std::{
42 collections::VecDeque,
43 fmt::Debug,
44 path::Path,
45 pin::pin,
46 sync::{
47 Arc,
48 atomic::{AtomicU8, Ordering},
49 },
50 time::Duration,
51};
52
53use bytes::Bytes;
54use nautilus_core::string::secret::REDACTED;
55use nautilus_cryptography::providers::install_cryptographic_provider;
56use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
57use tokio_tungstenite::tungstenite::{Error, client::IntoClientRequest, stream::Mode};
58
59use super::{SocketConfig, TcpMessageHandler, TcpReader, TcpWriter, WriterCommand};
60use crate::{
61 SocketStateSink,
62 backoff::{
63 ExponentialBackoff, RECONNECT_STABILITY_THRESHOLD, ReconnectThrottle, wait_reconnect_delay,
64 },
65 dst,
66 error::{SendError, is_connection_drop_io_error},
67 logging::{log_task_aborted, log_task_started, log_task_stopped},
68 mode::{
69 ConnectionMode, ControllerLifecycle, ReadSessionFence, ReconnectOutcome,
70 ReconnectRequestOutcome,
71 },
72 net::TcpStream,
73 tls::{create_tls_config_from_certs_dir, tcp_tls},
74};
75
76const CONNECTION_STATE_CHECK_INTERVAL_MS: u64 = 10;
78const GRACEFUL_SHUTDOWN_DELAY_MS: u64 = 100;
79const GRACEFUL_SHUTDOWN_TIMEOUT_SECS: u64 = 5;
80const WRITE_TIMEOUT_SECS: u64 = 5;
81
82const MAX_READ_BUFFER_BYTES: usize = 10 * 1024 * 1024;
84
85struct BufferedWrite {
86 data: Bytes,
87 replay_key: Option<u64>,
88}
89
90pub type SocketReconnectReplay = Arc<dyn Fn() -> Vec<Bytes> + Send + Sync>;
92
93struct SocketClientInner {
94 config: SocketConfig,
95 connector: Option<Arc<rustls::ClientConfig>>,
96 read_task: tokio::task::JoinHandle<()>,
97 read_fence: ReadSessionFence,
98 write_task: tokio::task::JoinHandle<()>,
99 writer_tx: tokio::sync::mpsc::UnboundedSender<WriterCommand>,
100 heartbeat_task: Option<tokio::task::JoinHandle<()>>,
101 connection_mode: Arc<AtomicU8>,
102 state_notify: Arc<tokio::sync::Notify>,
103 connect_timeout: Duration,
104 backoff: ExponentialBackoff,
105 reconnect_throttle: ReconnectThrottle,
106 reconnect_max_attempts: Option<u32>,
107 reconnect_attempt_count: u32,
108 state_sink: Option<SocketStateSink>,
109}
110
111impl SocketClientInner {
112 async fn connect_url(
118 config: SocketConfig,
119 state_sink: Option<SocketStateSink>,
120 ) -> anyhow::Result<Self> {
121 install_cryptographic_provider();
122
123 if config.suffix.is_empty() {
125 anyhow::bail!("Socket suffix cannot be empty: suffix is required for message framing");
126 }
127
128 config.validate()?;
131
132 let connect_timeout = Duration::from_millis(config.connect_timeout_ms.unwrap_or(10_000));
133 let reconnect_backoff = ExponentialBackoff::new(
134 Duration::from_millis(config.reconnect_delay_initial_ms.unwrap_or(2_000)),
135 Duration::from_millis(config.reconnect_delay_max_ms.unwrap_or(30_000)),
136 config.reconnect_backoff_factor.unwrap_or(1.5),
137 config.reconnect_jitter_ms.unwrap_or(100),
138 true, )?;
140 let connector = if let Some(dir) = &config.certs_dir {
141 let config = create_tls_config_from_certs_dir(Path::new(dir), false)?;
142 Some(Arc::new(config))
143 } else {
144 None
145 };
146
147 let max_retries = config.connection_max_retries.unwrap_or(5);
149
150 let mut backoff = ExponentialBackoff::new(
151 Duration::from_millis(500),
152 Duration::from_secs(5),
153 2.0,
154 250,
155 false,
156 )?;
157
158 let mut attempt = 0;
159 let (reader, writer) = loop {
160 attempt += 1;
161
162 let last_error = match dst::time::timeout(
163 connect_timeout,
164 Self::tls_connect_with_server(&config.url, config.mode, connector.clone()),
165 )
166 .await
167 {
168 Ok(Ok(result)) => {
169 if attempt > 1 {
170 log::info!("Socket connection established after {attempt} attempts");
171 }
172 break result;
173 }
174 Ok(Err(e)) => {
175 let error = e.to_string();
176 log::warn!(
177 "Socket connection attempt {attempt}/{max_retries} to {REDACTED} failed: {error}"
178 );
179 error
180 }
181 Err(_) => {
182 let error = format!(
183 "Connection timeout after {:.1}s (possible DNS resolution failure)",
184 connect_timeout.as_secs_f64()
185 );
186 log::warn!(
187 "Socket connection attempt {attempt}/{max_retries} to {REDACTED} timed out"
188 );
189 error
190 }
191 };
192
193 if attempt >= max_retries {
194 anyhow::bail!(
195 "Failed to connect to {REDACTED} after {max_retries} attempts: {last_error}. \
196 If this is a DNS error, check your network configuration and DNS settings."
197 );
198 }
199
200 let delay = backoff.next_duration();
201 log::debug!(
202 "Retrying in {delay:?} (attempt {}/{})",
203 attempt + 1,
204 max_retries
205 );
206 dst::time::sleep(delay).await;
207 };
208
209 log::debug!("Connected");
210
211 let connection_mode = Arc::new(AtomicU8::new(ConnectionMode::Reconnect.as_u8()));
212 let state_notify = Arc::new(tokio::sync::Notify::new());
213 let outcome =
214 ConnectionMode::complete_reconnect_with_sink(&connection_mode, state_sink.as_ref());
215 debug_assert_eq!(outcome, ReconnectOutcome::Reconnected);
216 let read_fence = ReadSessionFence::new();
217
218 let read_task = Self::spawn_read_task(
219 connection_mode.clone(),
220 read_fence.clone(),
221 reader,
222 config.message_handler.clone(),
223 config.suffix.clone(),
224 config.resolved_heartbeat_timeout(),
225 );
226
227 let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel::<WriterCommand>();
228
229 let write_task = Self::spawn_write_task(
230 connection_mode.clone(),
231 state_notify.clone(),
232 writer,
233 writer_rx,
234 config.suffix.clone(),
235 state_sink.clone(),
236 );
237
238 let heartbeat_task = config.heartbeat.as_ref().map(|heartbeat| {
239 Self::spawn_heartbeat_task(
240 connection_mode.clone(),
241 heartbeat.interval_secs,
242 heartbeat.payload.clone(),
243 writer_tx.clone(),
244 )
245 });
246 let reconnect_max_attempts = config.reconnect_max_attempts;
247
248 Ok(Self {
249 config,
250 connector,
251 read_task,
252 read_fence,
253 write_task,
254 writer_tx,
255 heartbeat_task,
256 connection_mode,
257 state_notify,
258 connect_timeout,
259 backoff: reconnect_backoff,
260 reconnect_throttle: ReconnectThrottle::default(),
261 reconnect_max_attempts,
262 reconnect_attempt_count: 0,
263 state_sink,
264 })
265 }
266
267 fn parse_socket_url(url: &str, mode: Mode) -> Result<(String, String), Error> {
277 if url.contains("://") {
278 let parsed = url.parse::<http::Uri>().map_err(|e| {
280 Error::Io(std::io::Error::new(
281 std::io::ErrorKind::InvalidInput,
282 format!("Invalid URL: {e}"),
283 ))
284 })?;
285
286 let host = parsed.host().ok_or_else(|| {
287 Error::Io(std::io::Error::new(
288 std::io::ErrorKind::InvalidInput,
289 "URL missing host",
290 ))
291 })?;
292
293 let port = parsed
294 .port_u16()
295 .unwrap_or_else(|| match parsed.scheme_str() {
296 Some("wss" | "https") => 443,
297 Some("ws" | "http") => 80,
298 _ => match mode {
299 Mode::Tls => 443,
300 Mode::Plain => 80,
301 },
302 });
303
304 Ok((format!("{host}:{port}"), url.to_string()))
305 } else {
306 let scheme = match mode {
309 Mode::Tls => "wss",
310 Mode::Plain => "ws",
311 };
312 Ok((url.to_string(), format!("{scheme}://{url}")))
313 }
314 }
315
316 pub(crate) async fn tls_connect_with_server(
326 url: &str,
327 mode: Mode,
328 connector: Option<Arc<rustls::ClientConfig>>,
329 ) -> Result<(TcpReader, TcpWriter), Error> {
330 log::debug!("Connecting to {REDACTED}");
331
332 let (socket_addr, request_url) = Self::parse_socket_url(url, mode)?;
333 let tcp_result = TcpStream::connect(&socket_addr).await;
334
335 match tcp_result {
336 Ok(stream) => {
337 log::debug!("TCP connection established to {socket_addr}, proceeding with TLS");
338
339 crate::net::apply_socket_options(&stream);
340
341 let request = request_url.into_client_request()?;
342 tcp_tls(&request, mode, stream, connector)
343 .await
344 .map(tokio::io::split)
345 }
346 Err(e) => {
347 log::warn!("TCP connection failed to {socket_addr}: {e:?}");
348 Err(Error::Io(e))
349 }
350 }
351 }
352
353 async fn reconnect(
364 &mut self,
365 reconnect_replay: Option<&SocketReconnectReplay>,
366 ) -> Result<ReconnectOutcome, Error> {
367 log::info!("Reconnecting");
368
369 if ConnectionMode::from_atomic(&self.connection_mode).is_disconnect() {
370 log::debug!("Reconnect aborted due to disconnect state");
371 return Ok(ReconnectOutcome::Aborted);
372 }
373
374 let (reader, new_writer) = dst::time::timeout(
376 self.connect_timeout,
377 Self::tls_connect_with_server(
378 &self.config.url,
379 self.config.mode,
380 self.connector.clone(),
381 ),
382 )
383 .await
384 .map_err(|_| {
385 Error::Io(std::io::Error::new(
386 std::io::ErrorKind::TimedOut,
387 format!(
388 "reconnection timed out after {}s",
389 self.connect_timeout.as_secs_f64()
390 ),
391 ))
392 })??;
393
394 if ConnectionMode::from_atomic(&self.connection_mode).is_disconnect() {
395 log::debug!("Reconnect aborted mid-flight (after connect)");
396 return Ok(ReconnectOutcome::Aborted);
397 }
398 log::debug!("Connected");
399
400 let (tx, rx) = tokio::sync::oneshot::channel();
404 let command = if let Some(reconnect_replay) = reconnect_replay {
405 WriterCommand::UpdateWithReplay(new_writer, reconnect_replay(), tx)
406 } else {
407 WriterCommand::Update(new_writer, tx)
408 };
409
410 if let Err(e) = self.writer_tx.send(command) {
411 log::error!("{e}");
412 return Err(Error::Io(std::io::Error::new(
413 std::io::ErrorKind::BrokenPipe,
414 format!("Failed to send update command: {e}"),
415 )));
416 }
417
418 match rx.await {
420 Ok(true) => log::debug!("Writer confirmed buffer drain success"),
421 Ok(false) => {
422 log::warn!("Writer failed to drain buffer, aborting reconnect");
423 return Err(Error::Io(std::io::Error::other(
425 "Failed to drain reconnection buffer",
426 )));
427 }
428 Err(e) => {
429 log::error!("Writer dropped update channel: {e}");
430 return Err(Error::Io(std::io::Error::new(
431 std::io::ErrorKind::BrokenPipe,
432 "Writer task dropped response channel",
433 )));
434 }
435 }
436
437 dst::time::sleep(Duration::from_millis(GRACEFUL_SHUTDOWN_DELAY_MS)).await;
439
440 if ConnectionMode::from_atomic(&self.connection_mode).is_disconnect() {
441 log::debug!("Reconnect aborted mid-flight (after delay)");
442 return Ok(ReconnectOutcome::Aborted);
443 }
444
445 self.read_fence.invalidate();
446
447 if !self.read_task.is_finished() {
448 self.read_task.abort();
449 log_task_aborted("read");
450 }
451
452 if ConnectionMode::complete_reconnect_with_sink(
455 &self.connection_mode,
456 self.state_sink.as_ref(),
457 ) == ReconnectOutcome::Aborted
458 {
459 log::debug!("Reconnect aborted (state changed during reconnect)");
460 return Ok(ReconnectOutcome::Aborted);
461 }
462
463 self.read_fence = ReadSessionFence::new();
465 self.read_task = Self::spawn_read_task(
466 self.connection_mode.clone(),
467 self.read_fence.clone(),
468 reader,
469 self.config.message_handler.clone(),
470 self.config.suffix.clone(),
471 self.config.resolved_heartbeat_timeout(),
472 );
473
474 log::info!("Reconnect succeeded");
475 Ok(ReconnectOutcome::Reconnected)
476 }
477
478 #[inline]
484 #[must_use]
485 pub(crate) fn is_alive(&self) -> bool {
486 !self.read_task.is_finished() && !self.write_task.is_finished()
487 }
488
489 #[must_use]
490 fn spawn_read_task<R>(
491 connection_state: Arc<AtomicU8>,
492 read_fence: ReadSessionFence,
493 reader: R,
494 handler: Option<TcpMessageHandler>,
495 suffix: Vec<u8>,
496 heartbeat_timeout_secs: Option<u64>,
497 ) -> tokio::task::JoinHandle<()>
498 where
499 R: AsyncRead + Unpin + Send + 'static,
500 {
501 log_task_started("read");
502
503 let check_interval = Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS);
505 let heartbeat_timeout = heartbeat_timeout_secs.map(Duration::from_secs);
506
507 tokio::task::spawn(Self::run_read_loop(
508 connection_state,
509 read_fence,
510 reader,
511 handler,
512 suffix,
513 heartbeat_timeout,
514 check_interval,
515 ))
516 }
517
518 async fn run_read_loop<R>(
519 connection_state: Arc<AtomicU8>,
520 read_fence: ReadSessionFence,
521 mut reader: R,
522 handler: Option<TcpMessageHandler>,
523 suffix: Vec<u8>,
524 heartbeat_timeout: Option<Duration>,
525 check_interval: Duration,
526 ) where
527 R: AsyncRead + Unpin,
528 {
529 let mut buf = Vec::new();
530 let mut last_data_time = dst::time::Instant::now();
531
532 'read: loop {
533 if !ConnectionMode::from_atomic(&connection_state).is_active() || !read_fence.is_valid()
534 {
535 if !buf.is_empty() {
536 log::debug!(
537 "Dropping {} buffered bytes after socket session ended",
538 buf.len()
539 );
540 buf.clear();
541 }
542 break;
543 }
544
545 match dst::time::timeout(check_interval, reader.read_buf(&mut buf)).await {
546 Ok(Ok(0)) => {
548 log::debug!("Connection closed by server");
549 break;
550 }
551 Ok(Err(e)) => {
552 log::debug!("Connection ended: {e}");
553 break;
554 }
555 Ok(Ok(bytes)) => {
557 log::trace!("Received <binary> {bytes} bytes");
558 last_data_time = dst::time::Instant::now();
559
560 if !ConnectionMode::from_atomic(&connection_state).is_active()
561 || !read_fence.is_valid()
562 {
563 log::debug!(
564 "Dropping {} buffered bytes after socket session ended",
565 buf.len()
566 );
567 buf.clear();
568 break;
569 }
570
571 while let Some((i, _)) = &buf
572 .windows(suffix.len())
573 .enumerate()
574 .find(|(_, pair)| pair.eq(&suffix))
575 {
576 let mut data: Vec<u8> = buf.drain(0..i + suffix.len()).collect();
577 data.truncate(data.len() - suffix.len());
578
579 if let Some(ref handler) = handler {
580 if !ConnectionMode::from_atomic(&connection_state).is_active()
581 || !read_fence.is_valid()
582 {
583 log::debug!(
584 "Dropping {} buffered bytes after socket session ended",
585 data.len() + buf.len()
586 );
587 buf.clear();
588 break 'read;
589 }
590 handler(&data);
591 }
592 }
593
594 if buf.len() > MAX_READ_BUFFER_BYTES {
595 log::error!(
596 "Read buffer exceeded maximum size ({MAX_READ_BUFFER_BYTES} bytes), closing connection"
597 );
598 break;
599 }
600 }
601 Err(_) => {
602 if let Some(timeout) = heartbeat_timeout {
603 let silent_for = last_data_time.elapsed();
604 if silent_for >= timeout {
605 log::warn!(
606 "Heartbeat timeout: no bytes received for {:.1}s",
607 silent_for.as_secs_f64()
608 );
609 break;
610 }
611 }
612 }
613 }
614 }
615
616 log_task_stopped("read");
617 }
618
619 async fn drain_reconnect_buffer<W>(
629 buffer: &mut VecDeque<BufferedWrite>,
630 writer: &mut W,
631 suffix: &[u8],
632 replay: &[Bytes],
633 ) -> bool
634 where
635 W: AsyncWrite + Unpin,
636 {
637 if buffer.is_empty() {
638 return false;
639 }
640
641 let initial_buffer_len = buffer.len();
642 log::info!("Sending {initial_buffer_len} buffered messages after reconnection");
643
644 while let Some(buffered) = buffer.front() {
645 if buffered.replay_key.is_some()
646 && replay.iter().any(|replayed| replayed == &buffered.data)
647 {
648 buffer.pop_front();
649 continue;
650 }
651
652 let mut combined_msg = Vec::with_capacity(buffered.data.len() + suffix.len());
653 combined_msg.extend_from_slice(&buffered.data);
654 combined_msg.extend_from_slice(suffix);
655
656 if let Err(e) = writer.write_all(&combined_msg).await {
657 if is_connection_drop_io_error(&e) {
658 log::warn!(
659 "Failed to send buffered message with suffix after reconnection: {e}, {} messages remain in buffer",
660 buffer.len()
661 );
662 } else {
663 log::error!(
664 "Failed to send buffered message with suffix after reconnection: {e}, {} messages remain in buffer",
665 buffer.len()
666 );
667 }
668 return true;
669 }
670
671 buffer.pop_front();
672 }
673
674 log::info!("Successfully sent all {initial_buffer_len} buffered messages");
675
676 false
677 }
678
679 async fn replace_writer<W>(
680 active_writer: &mut W,
681 new_writer: W,
682 replay: Vec<Bytes>,
683 reconnect_buffer: &mut VecDeque<BufferedWrite>,
684 suffix: &[u8],
685 ) -> bool
686 where
687 W: AsyncWrite + Unpin,
688 {
689 log::debug!("Received new writer");
690 dst::time::sleep(Duration::from_millis(GRACEFUL_SHUTDOWN_DELAY_MS)).await;
691
692 _ = dst::time::timeout(
693 Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS),
694 active_writer.shutdown(),
695 )
696 .await;
697
698 *active_writer = new_writer;
699 log::debug!("Updated writer");
700
701 let drain_result =
702 dst::time::timeout(Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS), async {
703 for replay_msg in &replay {
704 let mut framed = Vec::with_capacity(replay_msg.len() + suffix.len());
705 framed.extend_from_slice(replay_msg);
706 framed.extend_from_slice(suffix);
707 if let Err(e) = active_writer.write_all(&framed).await {
708 log::warn!("Failed to send reconnect replay: {e}");
709 return true;
710 }
711 }
712
713 Self::drain_reconnect_buffer(reconnect_buffer, active_writer, suffix, &replay).await
714 })
715 .await;
716
717 let send_error = drain_result.unwrap_or_else(|_| {
718 log::warn!(
719 "Timed out sending reconnect replay and buffered messages, {} buffered messages remain",
720 reconnect_buffer.len()
721 );
722 true
723 });
724 !send_error
725 }
726
727 fn spawn_write_task<W>(
728 connection_state: Arc<AtomicU8>,
729 state_notify: Arc<tokio::sync::Notify>,
730 writer: W,
731 mut writer_rx: tokio::sync::mpsc::UnboundedReceiver<WriterCommand<W>>,
732 suffix: Vec<u8>,
733 state_sink: Option<SocketStateSink>,
734 ) -> tokio::task::JoinHandle<()>
735 where
736 W: AsyncWrite + Unpin + Send + 'static,
737 {
738 log_task_started("write");
739
740 let check_interval = Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS);
742
743 tokio::task::spawn(async move {
744 let mut active_writer = writer;
745 let mut reconnect_buffer: VecDeque<BufferedWrite> = VecDeque::new();
746 let mut write_buf: Vec<u8> = Vec::new();
747
748 loop {
749 let mode = ConnectionMode::from_atomic(&connection_state);
750 if matches!(mode, ConnectionMode::Disconnect | ConnectionMode::Closed) {
751 break;
752 }
753
754 if mode.is_active() && !reconnect_buffer.is_empty() {
755 let drain_result = dst::time::timeout(
756 Duration::from_secs(WRITE_TIMEOUT_SECS),
757 Self::drain_reconnect_buffer(
758 &mut reconnect_buffer,
759 &mut active_writer,
760 &suffix,
761 &[],
762 ),
763 )
764 .await;
765 let send_error = drain_result.unwrap_or_else(|_| {
766 log::warn!(
767 "Timed out draining reconnect buffer after {WRITE_TIMEOUT_SECS}s, {} messages remain",
768 reconnect_buffer.len()
769 );
770 true
771 });
772
773 if send_error
774 && ConnectionMode::request_reconnect_with_sink(
775 &connection_state,
776 state_sink.as_ref(),
777 )
778 {
779 log::warn!("Writer triggering reconnect");
780 state_notify.notify_one();
781 }
782 continue;
783 }
784
785 match dst::time::timeout(check_interval, writer_rx.recv()).await {
786 Ok(Some(msg)) => {
787 let mode = ConnectionMode::from_atomic(&connection_state);
789 if matches!(mode, ConnectionMode::Disconnect | ConnectionMode::Closed) {
790 break;
791 }
792
793 match msg {
794 WriterCommand::Update(new_writer, tx) => {
795 let sent = Self::replace_writer(
796 &mut active_writer,
797 new_writer,
798 Vec::new(),
799 &mut reconnect_buffer,
800 &suffix,
801 )
802 .await;
803
804 if let Err(e) = tx.send(sent) {
805 log::error!(
806 "Failed to report drain status to controller: {e:?}"
807 );
808 }
809 }
810 WriterCommand::UpdateWithReplay(new_writer, replay, tx) => {
811 let sent = Self::replace_writer(
812 &mut active_writer,
813 new_writer,
814 replay,
815 &mut reconnect_buffer,
816 &suffix,
817 )
818 .await;
819
820 if let Err(e) = tx.send(sent) {
821 log::error!(
822 "Failed to report drain status to controller: {e:?}"
823 );
824 }
825 }
826 WriterCommand::Send(data)
827 if mode.is_reconnect() || !reconnect_buffer.is_empty() =>
828 {
829 log::debug!(
830 "Buffering message until reconnect drain completes ({} bytes)",
831 data.len()
832 );
833 Self::buffer_reconnect_write(&mut reconnect_buffer, data, None);
834 }
835 WriterCommand::SendOrReplay { key, data } if mode.is_reconnect() => {
836 log::debug!(
837 "Buffering replayable message while reconnecting ({} bytes)",
838 data.len()
839 );
840 Self::buffer_reconnect_write(
841 &mut reconnect_buffer,
842 data,
843 Some(key),
844 );
845 }
846 command @ (WriterCommand::Send(_)
847 | WriterCommand::SendOrReplay { .. }) => {
848 let (msg, replay_key) = match command {
849 WriterCommand::Send(data) => (data, None),
850 WriterCommand::SendOrReplay { key, data } => (data, Some(key)),
851 _ => unreachable!(),
852 };
853 write_buf.clear();
854 write_buf.extend_from_slice(&msg);
855 write_buf.extend_from_slice(&suffix);
856
857 let write_result = dst::time::timeout(
858 Duration::from_secs(WRITE_TIMEOUT_SECS),
859 active_writer.write_all(&write_buf),
860 )
861 .await;
862 let write_failed = match write_result {
863 Ok(Ok(())) => false,
864 Ok(Err(e)) => {
865 if is_connection_drop_io_error(&e) {
866 log::warn!("Failed to send message: {e}");
867 } else {
868 log::error!("Failed to send message: {e}");
869 }
870 true
871 }
872 Err(_) => {
873 log::warn!(
874 "Timed out sending message after {WRITE_TIMEOUT_SECS}s"
875 );
876 true
877 }
878 };
879
880 if write_failed {
881 Self::buffer_reconnect_write(
882 &mut reconnect_buffer,
883 msg,
884 replay_key,
885 );
886
887 if ConnectionMode::request_reconnect_with_sink(
889 &connection_state,
890 state_sink.as_ref(),
891 ) {
892 log::warn!("Writer triggering reconnect");
893 state_notify.notify_one();
894 }
895 }
896 }
897 }
898 }
899 Ok(None) => {
900 log::debug!("Writer channel closed, terminating writer task");
902 break;
903 }
904 Err(_) => {
905 }
907 }
908 }
909
910 _ = dst::time::timeout(
913 Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS),
914 active_writer.shutdown(),
915 )
916 .await;
917
918 log_task_stopped("write");
919 })
920 }
921
922 fn buffer_reconnect_write(
923 buffer: &mut VecDeque<BufferedWrite>,
924 data: Bytes,
925 replay_key: Option<u64>,
926 ) {
927 if let Some(key) = replay_key {
928 buffer.retain(|buffered| buffered.replay_key != Some(key));
929 }
930 buffer.push_back(BufferedWrite { data, replay_key });
931 }
932
933 fn spawn_heartbeat_task(
934 connection_state: Arc<AtomicU8>,
935 interval_secs: u64,
936 message: Vec<u8>,
937 writer_tx: tokio::sync::mpsc::UnboundedSender<WriterCommand>,
938 ) -> tokio::task::JoinHandle<()> {
939 log_task_started("heartbeat");
940
941 tokio::task::spawn(async move {
942 let interval = Duration::from_secs(interval_secs);
943
944 loop {
945 dst::time::sleep(interval).await;
946
947 match ConnectionMode::from_u8(connection_state.load(Ordering::SeqCst)) {
948 ConnectionMode::Active => {
949 let msg = WriterCommand::Send(message.clone().into());
950
951 match writer_tx.send(msg) {
952 Ok(()) => log::trace!("Sent heartbeat to writer task"),
953 Err(e) => {
954 log::warn!("Failed to send heartbeat to writer task: {e}");
955 }
956 }
957 }
958 ConnectionMode::Reconnect => {}
959 ConnectionMode::Disconnect | ConnectionMode::Closed => break,
960 }
961 }
962
963 log_task_stopped("heartbeat");
964 })
965 }
966}
967
968impl Drop for SocketClientInner {
969 fn drop(&mut self) {
970 self.read_fence.invalidate();
971
972 if !self.read_task.is_finished() {
973 self.read_task.abort();
974 log_task_aborted("read");
975 }
976
977 if !self.write_task.is_finished() {
978 self.write_task.abort();
979 log_task_aborted("write");
980 }
981
982 if let Some(ref handle) = self.heartbeat_task.take()
983 && !handle.is_finished()
984 {
985 handle.abort();
986 log_task_aborted("heartbeat");
987 }
988 }
989}
990
991pub struct SocketClient {
997 pub(crate) controller_task: tokio::task::JoinHandle<()>,
998 pub(crate) connection_mode: Arc<AtomicU8>,
999 pub(crate) state_notify: Arc<tokio::sync::Notify>,
1000 pub(crate) connect_timeout: Duration,
1001 pub writer_tx: tokio::sync::mpsc::UnboundedSender<WriterCommand>,
1002 state_sink: Option<SocketStateSink>,
1003 controller_lifecycle: Arc<ControllerLifecycle>,
1004 controller_notify: Arc<tokio::sync::Notify>,
1005}
1006
1007#[derive(Clone)]
1009pub struct SocketReconnectHandle {
1010 connection_mode: Arc<AtomicU8>,
1011 state_sink: Option<SocketStateSink>,
1012 controller_lifecycle: Arc<ControllerLifecycle>,
1013 controller_notify: Arc<tokio::sync::Notify>,
1014}
1015
1016impl Debug for SocketReconnectHandle {
1017 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1018 f.debug_struct(stringify!(SocketReconnectHandle))
1019 .field(
1020 "connection_mode",
1021 &ConnectionMode::from_atomic(&self.connection_mode),
1022 )
1023 .finish_non_exhaustive()
1024 }
1025}
1026
1027impl SocketReconnectHandle {
1028 #[must_use]
1033 pub fn request_reconnect(&self) -> ReconnectRequestOutcome {
1034 let Some(_request) = self.controller_lifecycle.enter_request() else {
1035 return ReconnectRequestOutcome::Closed;
1036 };
1037
1038 let outcome = ConnectionMode::request_reconnect_outcome_with_sink(
1039 &self.connection_mode,
1040 self.state_sink.as_ref(),
1041 );
1042
1043 if outcome == ReconnectRequestOutcome::Accepted {
1044 self.controller_notify.notify_one();
1045 }
1046 outcome
1047 }
1048}
1049
1050impl Debug for SocketClient {
1051 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1052 f.debug_struct(stringify!(SocketClient)).finish()
1053 }
1054}
1055
1056#[bon::bon]
1057impl SocketClient {
1058 #[builder(finish_fn = connect)]
1071 pub async fn builder(
1072 config: SocketConfig,
1073 post_reconnection: Option<Arc<dyn Fn() + Send + Sync>>,
1074 state_sink: Option<SocketStateSink>,
1075 reconnect_replay: Option<SocketReconnectReplay>,
1076 ) -> anyhow::Result<Self> {
1077 let inner = SocketClientInner::connect_url(config, state_sink).await?;
1078 let writer_tx = inner.writer_tx.clone();
1079 let connection_mode = inner.connection_mode.clone();
1080 let state_notify = inner.state_notify.clone();
1081 let connect_timeout = inner.connect_timeout;
1082 let state_sink = inner.state_sink.clone();
1083 let controller_lifecycle = Arc::new(ControllerLifecycle::new());
1084 let controller_notify = Arc::new(tokio::sync::Notify::new());
1085 let controller_task = Self::spawn_controller_task(
1086 inner,
1087 connection_mode.clone(),
1088 state_notify.clone(),
1089 Arc::clone(&controller_lifecycle),
1090 Arc::clone(&controller_notify),
1091 post_reconnection,
1092 reconnect_replay,
1093 );
1094 let abort_handle = controller_task.abort_handle();
1095 controller_lifecycle.set_abort(move || abort_handle.abort());
1096
1097 Ok(Self {
1098 controller_task,
1099 connection_mode,
1100 state_notify,
1101 connect_timeout,
1102 writer_tx,
1103 state_sink,
1104 controller_lifecycle,
1105 controller_notify,
1106 })
1107 }
1108
1109 #[must_use]
1111 pub fn reconnect_handle(&self) -> SocketReconnectHandle {
1112 SocketReconnectHandle {
1113 connection_mode: Arc::clone(&self.connection_mode),
1114 state_sink: self.state_sink.clone(),
1115 controller_lifecycle: Arc::clone(&self.controller_lifecycle),
1116 controller_notify: Arc::clone(&self.controller_notify),
1117 }
1118 }
1119
1120 #[must_use]
1125 pub fn request_reconnect(&self) -> bool {
1126 self.reconnect_handle().request_reconnect() == ReconnectRequestOutcome::Accepted
1127 }
1128
1129 #[must_use]
1131 pub fn connection_mode(&self) -> ConnectionMode {
1132 ConnectionMode::from_atomic(&self.connection_mode)
1133 }
1134
1135 #[inline]
1140 #[must_use]
1141 pub fn is_active(&self) -> bool {
1142 self.connection_mode().is_active()
1143 }
1144
1145 #[inline]
1150 #[must_use]
1151 pub fn is_reconnecting(&self) -> bool {
1152 self.connection_mode().is_reconnect()
1153 }
1154
1155 #[inline]
1159 #[must_use]
1160 pub fn is_disconnecting(&self) -> bool {
1161 self.connection_mode().is_disconnect()
1162 }
1163
1164 #[inline]
1170 #[must_use]
1171 pub fn is_closed(&self) -> bool {
1172 self.connection_mode().is_closed()
1173 }
1174
1175 pub async fn close(&self) {
1180 self.begin_shutdown();
1181 self.close_with_timeout(Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS))
1182 .await;
1183 }
1184
1185 pub fn begin_shutdown(&self) {
1187 ConnectionMode::request_disconnect(&self.connection_mode);
1188 self.state_notify.notify_waiters();
1189 }
1190
1191 async fn close_with_timeout(&self, shutdown_timeout: Duration) {
1192 self.begin_shutdown();
1193
1194 if dst::time::timeout(shutdown_timeout, async {
1195 while !self.controller_task.is_finished() {
1196 dst::time::sleep(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
1197 }
1198 })
1199 .await
1200 .is_err()
1201 {
1202 log::warn!("Timeout waiting for controller task to finish");
1203 }
1204
1205 if !self.controller_task.is_finished() {
1206 self.controller_task.abort();
1207 log_task_aborted("controller");
1208 }
1209
1210 self.connection_mode
1211 .store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
1212 self.state_notify.notify_waiters();
1213 }
1214
1215 #[inline]
1219 fn check_not_terminal(&self) -> Result<(), SendError> {
1220 match self.connection_mode() {
1221 ConnectionMode::Disconnect | ConnectionMode::Closed => Err(SendError::Closed),
1222 _ => Ok(()),
1223 }
1224 }
1225
1226 async fn wait_for_active(&self) -> Result<(), SendError> {
1232 const FALLBACK_INTERVAL_MS: u64 = 100;
1233
1234 let mode = self.connection_mode();
1235 if mode.is_active() {
1236 return Ok(());
1237 }
1238
1239 if matches!(mode, ConnectionMode::Disconnect | ConnectionMode::Closed) {
1240 return Err(SendError::Closed);
1241 }
1242
1243 log::debug!("Waiting for client to become ACTIVE before sending...");
1244
1245 let fallback_interval = Duration::from_millis(FALLBACK_INTERVAL_MS);
1246
1247 dst::time::timeout(self.connect_timeout, async {
1248 loop {
1249 let mut notified = pin!(self.state_notify.notified());
1251 notified.as_mut().enable();
1252
1253 let mode = self.connection_mode();
1254 if mode.is_active() {
1255 return Ok(());
1256 }
1257
1258 if matches!(mode, ConnectionMode::Disconnect | ConnectionMode::Closed) {
1259 return Err(());
1260 }
1261
1262 tokio::select! {
1263 biased;
1264 () = notified => {}
1265 () = dst::time::sleep(fallback_interval) => {}
1266 }
1267 }
1268 })
1269 .await
1270 .map_err(|_| SendError::Timeout)?
1271 .map_err(|()| SendError::Closed)
1272 }
1273
1274 pub async fn send_bytes(&self, data: Vec<u8>) -> Result<(), SendError> {
1284 self.check_not_terminal()?;
1285 self.wait_for_active().await?;
1286
1287 let msg = WriterCommand::Send(data.into());
1288 self.writer_tx
1289 .send(msg)
1290 .map_err(|e| SendError::BrokenPipe(e.to_string()))
1291 }
1292
1293 fn spawn_controller_task(
1294 mut inner: SocketClientInner,
1295 connection_mode: Arc<AtomicU8>,
1296 state_notify: Arc<tokio::sync::Notify>,
1297 controller_lifecycle: Arc<ControllerLifecycle>,
1298 controller_notify: Arc<tokio::sync::Notify>,
1299 post_reconnection: Option<Arc<dyn Fn() + Send + Sync>>,
1300 reconnect_replay: Option<SocketReconnectReplay>,
1301 ) -> tokio::task::JoinHandle<()> {
1302 const CONTROLLER_FALLBACK_INTERVAL_MS: u64 = 100;
1303
1304 tokio::task::spawn(async move {
1305 let _activity = controller_lifecycle.activity();
1306 log_task_started("controller");
1307
1308 let fallback_interval = Duration::from_millis(CONTROLLER_FALLBACK_INTERVAL_MS);
1309 let mut reconnected_at = None;
1310
1311 loop {
1312 tokio::select! {
1313 biased;
1314 () = controller_notify.notified() => {}
1315 () = state_notify.notified() => {}
1316 () = dst::time::sleep(fallback_interval) => {}
1317 }
1318
1319 let mut mode = ConnectionMode::from_atomic(&connection_mode);
1320
1321 if mode.is_disconnect() {
1322 log::debug!("Disconnecting");
1323
1324 let timeout = Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS);
1325 if dst::time::timeout(timeout, async {
1326 dst::time::sleep(Duration::from_millis(GRACEFUL_SHUTDOWN_DELAY_MS)).await;
1328
1329 inner.read_fence.invalidate();
1330 if !inner.read_task.is_finished() {
1331 inner.read_task.abort();
1332 log_task_aborted("read");
1333 }
1334
1335 if let Some(task) = &inner.heartbeat_task
1336 && !task.is_finished()
1337 {
1338 task.abort();
1339 log_task_aborted("heartbeat");
1340 }
1341 })
1342 .await
1343 .is_err()
1344 {
1345 log::warn!("Shutdown timed out after {}s", timeout.as_secs());
1346 }
1347
1348 log::debug!("Closed");
1349 connection_mode.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
1350 state_notify.notify_waiters();
1351 break; }
1353
1354 if mode.is_closed() {
1355 log::debug!("Connection closed");
1356
1357 inner.read_fence.invalidate();
1358 if !inner.read_task.is_finished() {
1359 inner.read_task.abort();
1360 log_task_aborted("read");
1361 }
1362
1363 if let Some(task) = &inner.heartbeat_task
1364 && !task.is_finished()
1365 {
1366 task.abort();
1367 log_task_aborted("heartbeat");
1368 }
1369
1370 state_notify.notify_waiters();
1371 break;
1372 }
1373
1374 if mode.is_active() && !inner.is_alive() {
1375 if ConnectionMode::request_reconnect_with_sink(
1376 &connection_mode,
1377 inner.state_sink.as_ref(),
1378 ) {
1379 log::info!("Detected dead read task, transitioning to RECONNECT");
1380 }
1381 mode = ConnectionMode::from_atomic(&connection_mode);
1382 }
1383
1384 if mode.is_reconnect() {
1385 let reconnect_uptime = reconnected_at
1386 .take()
1387 .map(|started: dst::time::Instant| started.elapsed());
1388 let previous_reconnect_stable = reconnect_uptime
1389 .is_some_and(|uptime| uptime >= RECONNECT_STABILITY_THRESHOLD);
1390
1391 if previous_reconnect_stable {
1392 inner.backoff.reset();
1393 inner.reconnect_attempt_count = 0;
1394 log::debug!(
1395 "Socket remained active for at least {}s, resetting reconnect cycle",
1396 RECONNECT_STABILITY_THRESHOLD.as_secs()
1397 );
1398 }
1399
1400 if let Some(max_attempts) = inner.reconnect_max_attempts
1402 && inner.reconnect_attempt_count >= max_attempts
1403 {
1404 log::error!(
1405 "Max reconnection attempts ({max_attempts}) exceeded, transitioning to CLOSED"
1406 );
1407
1408 if connection_mode
1409 .compare_exchange(
1410 ConnectionMode::Reconnect.as_u8(),
1411 ConnectionMode::Closed.as_u8(),
1412 Ordering::SeqCst,
1413 Ordering::SeqCst,
1414 )
1415 .is_ok()
1416 {
1417 state_notify.notify_waiters();
1418 break;
1419 }
1420 continue;
1421 }
1422
1423 let backoff_delay = if reconnect_uptime.is_some() && !previous_reconnect_stable
1424 {
1425 inner.backoff.next_duration()
1426 } else {
1427 Duration::ZERO
1428 };
1429
1430 let duration = inner.reconnect_throttle.gated_delay(backoff_delay);
1431 if !duration.is_zero() {
1432 log::warn!("Backing off for {}s...", duration.as_secs_f64());
1433
1434 if !wait_reconnect_delay(
1435 duration,
1436 connection_mode.as_ref(),
1437 state_notify.as_ref(),
1438 )
1439 .await
1440 {
1441 log::debug!("Backoff interrupted by terminal state");
1442 continue;
1443 }
1444 }
1445
1446 inner.reconnect_attempt_count += 1;
1447 inner.reconnect_throttle.record_attempt();
1448
1449 let reconnect_result = tokio::select! {
1451 biased;
1452 result = inner.reconnect(reconnect_replay.as_ref()) => Some(result),
1453 () = async {
1454 loop {
1455 let mut notified = pin!(state_notify.notified());
1457 notified.as_mut().enable();
1458
1459 if ConnectionMode::from_atomic(&connection_mode).is_disconnect() {
1460 break;
1461 }
1462 notified.await;
1463 }
1464 } => None,
1465 };
1466
1467 match reconnect_result {
1468 None => {
1469 log::debug!("Reconnect interrupted by disconnect");
1470 }
1471 Some(Ok(ReconnectOutcome::Reconnected)) => {
1472 log::debug!("Reconnected successfully");
1473 reconnected_at = Some(dst::time::Instant::now());
1474
1475 state_notify.notify_waiters();
1476
1477 if ConnectionMode::from_atomic(&connection_mode).is_active() {
1481 if let Some(ref handler) = post_reconnection {
1482 handler();
1483 log::debug!("Called `post_reconnection` handler");
1484 }
1485 } else {
1486 log::debug!(
1487 "Skipping post_reconnection handlers due to disconnect state"
1488 );
1489 }
1490 }
1491 Some(Ok(ReconnectOutcome::Aborted)) => {
1492 log::debug!("Reconnect aborted");
1493 }
1494 Some(Err(e)) => {
1495 let duration = inner.backoff.next_duration();
1496 log::warn!(
1497 "Reconnect attempt {} failed: {e}",
1498 inner.reconnect_attempt_count
1499 );
1500
1501 if !duration.is_zero() {
1502 log::warn!("Backing off for {}s...", duration.as_secs_f64());
1503 if !wait_reconnect_delay(
1504 duration,
1505 connection_mode.as_ref(),
1506 state_notify.as_ref(),
1507 )
1508 .await
1509 {
1510 log::debug!("Backoff interrupted by terminal state");
1511 }
1512 }
1513 }
1514 }
1515 }
1516 }
1517 log_task_stopped("controller");
1518 })
1519 }
1520}
1521
1522impl Drop for SocketClient {
1524 fn drop(&mut self) {
1525 let controller_running = !self.controller_task.is_finished();
1526 self.controller_lifecycle.close_and_abort();
1527
1528 if controller_running {
1529 log_task_aborted("controller");
1530 }
1531 }
1532}
1533
1534#[cfg(test)]
1535#[cfg(not(feature = "turmoil"))]
1536#[cfg(not(all(feature = "simulation", madsim)))] #[cfg(target_os = "linux")] mod tests {
1539 use log::Level;
1540 use nautilus_common::testing::wait_until_async;
1541 use parking_lot::Mutex as BlockingMutex;
1542 use rstest::rstest;
1543 use tokio::{
1544 io::{AsyncReadExt, AsyncWriteExt},
1545 net::{TcpListener, TcpStream},
1546 sync::Mutex,
1547 task,
1548 time::{Duration, sleep},
1549 };
1550
1551 use super::*;
1552 use crate::{SocketState, logging::tests::capture_logs_for, socket::SocketHeartbeat};
1553
1554 async fn bind_test_server() -> (u16, TcpListener) {
1555 let listener = TcpListener::bind("127.0.0.1:0")
1556 .await
1557 .expect("Failed to bind ephemeral port");
1558 let port = listener.local_addr().unwrap().port();
1559 (port, listener)
1560 }
1561
1562 async fn run_echo_server(mut socket: TcpStream) {
1563 let mut buf = Vec::new();
1564 loop {
1565 match socket.read_buf(&mut buf).await {
1566 Ok(0) => {
1567 break;
1568 }
1569 Ok(_n) => {
1570 while let Some(idx) = buf.array_windows().position(|w| w == b"\r\n") {
1571 let mut line = buf.drain(..idx + 2).collect::<Vec<u8>>();
1572 line.truncate(line.len() - 2);
1574
1575 if line == b"close" {
1576 let _ = socket.shutdown().await;
1577 return;
1578 }
1579
1580 let mut echo_data = line;
1581 echo_data.extend_from_slice(b"\r\n");
1582 if socket.write_all(&echo_data).await.is_err() {
1583 break;
1584 }
1585 }
1586 }
1587 Err(e) => {
1588 eprintln!("Server read error: {e}");
1589 break;
1590 }
1591 }
1592 }
1593 }
1594
1595 #[rstest]
1596 #[case::drain_failed(
1597 Some(false),
1598 std::io::ErrorKind::Other,
1599 "Failed to drain reconnection buffer"
1600 )]
1601 #[case::reply_dropped(
1602 None,
1603 std::io::ErrorKind::BrokenPipe,
1604 "Writer task dropped response channel"
1605 )]
1606 #[tokio::test]
1607 async fn test_reconnect_writer_failure_preserves_reconnect_state(
1608 #[case] reply: Option<bool>,
1609 #[case] kind: std::io::ErrorKind,
1610 #[case] message: &str,
1611 ) {
1612 let (port, listener) = bind_test_server().await;
1613 let config = SocketConfig::builder()
1614 .url(format!("127.0.0.1:{port}"))
1615 .mode(Mode::Plain)
1616 .suffix(b"\r\n".to_vec())
1617 .build()
1618 .unwrap();
1619 let mut client = SocketClientInner::connect_url(config, None).await.unwrap();
1620 let (_peer, _) = listener.accept().await.unwrap();
1621 client.read_task.abort();
1622 client.write_task.abort();
1623 let _ = (&mut client.read_task).await;
1624 let _ = (&mut client.write_task).await;
1625 let (writer_tx, mut writer_rx) = tokio::sync::mpsc::unbounded_channel();
1626 client.writer_tx = writer_tx;
1627 client
1628 .connection_mode
1629 .store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
1630
1631 let writer = tokio::spawn(async move {
1632 let WriterCommand::Update(_, sender) = writer_rx.recv().await.unwrap() else {
1633 panic!("expected writer update");
1634 };
1635
1636 if let Some(reply) = reply {
1637 sender.send(reply).unwrap();
1638 }
1639 });
1640
1641 let error = client.reconnect(None).await.unwrap_err();
1642 writer.await.unwrap();
1643
1644 let Error::Io(error) = error else {
1645 panic!("expected I/O error, was {error}");
1646 };
1647
1648 assert_eq!(error.kind(), kind);
1649 assert_eq!(error.to_string(), message);
1650 assert_eq!(
1651 ConnectionMode::from_atomic(&client.connection_mode),
1652 ConnectionMode::Reconnect
1653 );
1654 }
1655
1656 #[tokio::test]
1657 async fn test_basic_send_receive() {
1658 let (port, listener) = bind_test_server().await;
1659 let server_task = task::spawn(async move {
1660 let (socket, _) = listener.accept().await.unwrap();
1661 run_echo_server(socket).await;
1662 });
1663
1664 let config = SocketConfig {
1665 url: format!("127.0.0.1:{port}"),
1666 mode: Mode::Plain,
1667 suffix: b"\r\n".to_vec(),
1668 message_handler: None,
1669 heartbeat: None,
1670 connect_timeout_ms: None,
1671 reconnect_delay_initial_ms: None,
1672 reconnect_backoff_factor: None,
1673 reconnect_delay_max_ms: None,
1674 reconnect_jitter_ms: None,
1675 reconnect_max_attempts: None,
1676 connection_max_retries: None,
1677 heartbeat_timeout_secs: None,
1678 certs_dir: None,
1679 };
1680
1681 let client = SocketClient::builder()
1682 .config(config)
1683 .connect()
1684 .await
1685 .expect("Client connect failed unexpectedly");
1686
1687 client.send_bytes(b"Hello".into()).await.unwrap();
1688 client.send_bytes(b"World".into()).await.unwrap();
1689
1690 sleep(Duration::from_millis(100)).await;
1692
1693 client.send_bytes(b"close".into()).await.unwrap();
1694 server_task.await.unwrap();
1695 assert!(!client.is_closed());
1696 }
1697
1698 #[tokio::test]
1699 async fn test_reconnect_fail_exhausted() {
1700 let (port, listener) = bind_test_server().await;
1701 drop(listener); wait_until_async(
1705 || async {
1706 TcpStream::connect(format!("127.0.0.1:{port}"))
1707 .await
1708 .is_err()
1709 },
1710 Duration::from_secs(2),
1711 )
1712 .await;
1713
1714 let config = SocketConfig {
1715 url: format!("127.0.0.1:{port}"),
1716 mode: Mode::Plain,
1717 suffix: b"\r\n".to_vec(),
1718 message_handler: None,
1719 heartbeat: None,
1720 connect_timeout_ms: Some(100),
1721 reconnect_delay_initial_ms: Some(50),
1722 reconnect_backoff_factor: Some(1.0),
1723 reconnect_delay_max_ms: Some(50),
1724 reconnect_jitter_ms: Some(0),
1725 connection_max_retries: Some(1),
1726 reconnect_max_attempts: None,
1727 heartbeat_timeout_secs: None,
1728 certs_dir: None,
1729 };
1730
1731 let client_res = SocketClient::builder().config(config).connect().await;
1732 assert!(
1733 client_res.is_err(),
1734 "Should fail quickly with no server listening"
1735 );
1736 }
1737
1738 #[rstest]
1739 #[tokio::test]
1740 async fn test_connect_failure_redacts_url_from_logs_and_error() {
1741 const ENDPOINT_PATH_SECRET: &str = "unique-endpoint-path-secret";
1742 const ENDPOINT_QUERY_SECRET: &str = "unique-endpoint-query-secret";
1743
1744 let (port, listener) = bind_test_server().await;
1745 drop(listener); wait_until_async(
1749 || async {
1750 TcpStream::connect(format!("127.0.0.1:{port}"))
1751 .await
1752 .is_err()
1753 },
1754 Duration::from_secs(2),
1755 )
1756 .await;
1757
1758 let capture = capture_logs_for(&["nautilus_network::socket::client"]).await;
1759 let config = SocketConfig {
1760 url: format!(
1761 "ws://127.0.0.1:{port}/{ENDPOINT_PATH_SECRET}?api_key={ENDPOINT_QUERY_SECRET}"
1762 ),
1763 mode: Mode::Plain,
1764 suffix: b"\r\n".to_vec(),
1765 message_handler: None,
1766 heartbeat: None,
1767 connect_timeout_ms: Some(1_000),
1768 reconnect_delay_initial_ms: Some(50),
1769 reconnect_delay_max_ms: Some(50),
1770 reconnect_backoff_factor: Some(1.0),
1771 reconnect_jitter_ms: Some(0),
1772 connection_max_retries: Some(1),
1773 reconnect_max_attempts: None,
1774 heartbeat_timeout_secs: None,
1775 certs_dir: None,
1776 };
1777
1778 let error = SocketClient::builder()
1779 .config(config)
1780 .connect()
1781 .await
1782 .expect_err("connection should fail with no server listening");
1783
1784 let error_message = error.to_string();
1785 assert!(
1786 error_message.starts_with(&format!("Failed to connect to {REDACTED}")),
1787 "connect error omitted the redaction placeholder: {error_message}"
1788 );
1789 assert!(
1790 !error_message.contains(ENDPOINT_PATH_SECRET)
1791 && !error_message.contains(ENDPOINT_QUERY_SECRET),
1792 "connect error exposed the secret marker: {error_message}"
1793 );
1794
1795 let messages = capture.messages();
1796 let (_, attempt_warning) = messages
1799 .iter()
1800 .find(|(level, message)| {
1801 *level == Level::Warn && message.starts_with("Socket connection attempt 1/1 ")
1802 })
1803 .expect("connect-failure warning was not captured");
1804 assert!(
1805 attempt_warning.contains(REDACTED),
1806 "connect-failure warning omitted the redaction placeholder: {attempt_warning}"
1807 );
1808 assert!(
1809 messages.iter().any(|(level, message)| {
1810 *level == Level::Debug && *message == format!("Connecting to {REDACTED}")
1811 }),
1812 "connect debug log should redact the endpoint, was {messages:?}"
1813 );
1814 assert!(
1815 messages.iter().all(|(_, message)| {
1816 !message.contains(ENDPOINT_PATH_SECRET) && !message.contains(ENDPOINT_QUERY_SECRET)
1817 }),
1818 "socket client logs exposed the secret marker: {messages:?}"
1819 );
1820 }
1821
1822 #[tokio::test]
1823 async fn test_user_disconnect() {
1824 let (port, listener) = bind_test_server().await;
1825 let server_task = task::spawn(async move {
1826 let (socket, _) = listener.accept().await.unwrap();
1827 let mut buf = [0u8; 1024];
1828 let _ = socket.try_read(&mut buf);
1829
1830 loop {
1831 sleep(Duration::from_secs(1)).await;
1832 }
1833 });
1834
1835 let config = SocketConfig {
1836 url: format!("127.0.0.1:{port}"),
1837 mode: Mode::Plain,
1838 suffix: b"\r\n".to_vec(),
1839 message_handler: None,
1840 heartbeat: None,
1841 connect_timeout_ms: None,
1842 reconnect_delay_initial_ms: None,
1843 reconnect_backoff_factor: None,
1844 reconnect_delay_max_ms: None,
1845 reconnect_jitter_ms: None,
1846 reconnect_max_attempts: None,
1847 connection_max_retries: None,
1848 heartbeat_timeout_secs: None,
1849 certs_dir: None,
1850 };
1851
1852 let client = SocketClient::builder()
1853 .config(config)
1854 .connect()
1855 .await
1856 .unwrap();
1857
1858 client.close().await;
1859 assert!(client.is_closed());
1860 server_task.abort();
1861 }
1862
1863 #[tokio::test]
1864 async fn test_close_after_closed_returns_fast_and_preserves_state() {
1865 let (port, listener) = bind_test_server().await;
1866
1867 let server_task = task::spawn(async move {
1868 let (socket, _) = listener.accept().await.unwrap();
1871 drop(socket);
1872 drop(listener);
1873 sleep(Duration::from_secs(5)).await;
1874 });
1875
1876 let config = SocketConfig {
1877 url: format!("127.0.0.1:{port}"),
1878 mode: Mode::Plain,
1879 suffix: b"\r\n".to_vec(),
1880 message_handler: None,
1881 heartbeat: None,
1882 connect_timeout_ms: Some(200),
1883 reconnect_delay_initial_ms: Some(50),
1884 reconnect_backoff_factor: Some(1.0),
1885 reconnect_delay_max_ms: Some(50),
1886 reconnect_jitter_ms: Some(0),
1887 connection_max_retries: None,
1888 reconnect_max_attempts: Some(1),
1889 heartbeat_timeout_secs: None,
1890 certs_dir: None,
1891 };
1892
1893 let client = SocketClient::builder()
1894 .config(config)
1895 .connect()
1896 .await
1897 .unwrap();
1898
1899 wait_until_async(|| async { client.is_closed() }, Duration::from_secs(5)).await;
1900
1901 let start = std::time::Instant::now();
1905 client.close().await;
1906 let elapsed = start.elapsed();
1907
1908 assert!(client.is_closed(), "Client should remain CLOSED");
1909 assert!(
1910 !client.is_disconnecting(),
1911 "Closed client should not report DISCONNECT after close()"
1912 );
1913 assert!(
1914 elapsed < Duration::from_secs(2),
1915 "close() on a closed client should return fast, took {elapsed:?}"
1916 );
1917
1918 server_task.abort();
1919 }
1920
1921 #[rstest]
1922 #[tokio::test]
1923 async fn test_heartbeat_task_warns_when_writer_channel_closed() {
1924 let capture = capture_logs_for(&["nautilus_network::socket::client"]).await;
1925 let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
1926 let (writer_tx, mut writer_rx) = tokio::sync::mpsc::unbounded_channel();
1927 let task = SocketClientInner::spawn_heartbeat_task(
1928 Arc::clone(&connection_state),
1929 1,
1930 b"ping".to_vec(),
1931 writer_tx,
1932 );
1933
1934 tokio::time::timeout(Duration::from_secs(2), writer_rx.recv())
1935 .await
1936 .expect("timed out waiting for the first heartbeat")
1937 .expect("heartbeat channel closed before the first heartbeat");
1938 drop(writer_rx);
1939 tokio::time::sleep(Duration::from_millis(1_200)).await;
1940
1941 let messages = capture.messages();
1942 assert!(
1943 messages.iter().any(|(level, message)| {
1944 *level == Level::Warn && message.contains("Failed to send heartbeat to writer task")
1945 }),
1946 "expected WARN when the writer channel is closed, was {messages:?}"
1947 );
1948 assert!(
1949 !messages.iter().any(|(level, _)| *level == Level::Error),
1950 "a closed writer channel must not log ERROR, was {messages:?}"
1951 );
1952
1953 connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
1954 tokio::time::timeout(Duration::from_secs(2), task)
1955 .await
1956 .expect("heartbeat task should stop after close")
1957 .unwrap();
1958 }
1959
1960 #[tokio::test]
1961 async fn test_heartbeat() {
1962 let (port, listener) = bind_test_server().await;
1963 let received = Arc::new(Mutex::new(Vec::new()));
1964 let received2 = received.clone();
1965
1966 let server_task = task::spawn(async move {
1967 let (socket, _) = listener.accept().await.unwrap();
1968
1969 let mut buf = Vec::new();
1970 loop {
1971 match socket.try_read_buf(&mut buf) {
1972 Ok(0) => break,
1973 Ok(_) => {
1974 while let Some(idx) = buf.array_windows().position(|w| w == b"\r\n") {
1975 let mut line = buf.drain(..idx + 2).collect::<Vec<u8>>();
1976 line.truncate(line.len() - 2);
1977 received2.lock().await.push(line);
1978 }
1979 }
1980 Err(_) => {
1981 tokio::time::sleep(Duration::from_millis(10)).await;
1982 }
1983 }
1984 }
1985 });
1986
1987 let config = SocketConfig {
1988 url: format!("127.0.0.1:{port}"),
1989 mode: Mode::Plain,
1990 suffix: b"\r\n".to_vec(),
1991 message_handler: None,
1992 heartbeat: Some(SocketHeartbeat {
1993 interval_secs: 1,
1994 payload: b"ping".to_vec(),
1995 }),
1996 connect_timeout_ms: None,
1997 reconnect_delay_initial_ms: None,
1998 reconnect_backoff_factor: None,
1999 reconnect_delay_max_ms: None,
2000 reconnect_jitter_ms: None,
2001 reconnect_max_attempts: None,
2002 connection_max_retries: None,
2003 heartbeat_timeout_secs: None,
2004 certs_dir: None,
2005 };
2006
2007 let client = SocketClient::builder()
2008 .config(config)
2009 .connect()
2010 .await
2011 .unwrap();
2012
2013 sleep(Duration::from_secs(3)).await;
2015
2016 {
2017 let lock = received.lock().await;
2018 let pings = lock
2019 .iter()
2020 .filter(|line| line == &&b"ping".to_vec())
2021 .count();
2022 assert!(
2023 pings >= 2,
2024 "Expected at least 2 heartbeat pings; got {pings}"
2025 );
2026 }
2027
2028 client.close().await;
2029 server_task.abort();
2030 }
2031
2032 #[tokio::test]
2033 async fn test_reconnect_success() {
2034 let (port, listener) = bind_test_server().await;
2035
2036 let server_task = task::spawn(async move {
2040 let (mut socket, _) = listener.accept().await.expect("First accept failed");
2042
2043 sleep(Duration::from_millis(500)).await;
2045 let _ = socket.shutdown().await;
2046
2047 sleep(Duration::from_millis(500)).await;
2049
2050 let (socket, _) = listener.accept().await.expect("Second accept failed");
2052 run_echo_server(socket).await;
2053 });
2054
2055 let config = SocketConfig {
2056 url: format!("127.0.0.1:{port}"),
2057 mode: Mode::Plain,
2058 suffix: b"\r\n".to_vec(),
2059 message_handler: None,
2060 heartbeat: None,
2061 connect_timeout_ms: Some(5_000),
2062 reconnect_delay_initial_ms: Some(500),
2063 reconnect_delay_max_ms: Some(5_000),
2064 reconnect_backoff_factor: Some(2.0),
2065 reconnect_jitter_ms: Some(50),
2066 reconnect_max_attempts: None,
2067 connection_max_retries: None,
2068 heartbeat_timeout_secs: None,
2069 certs_dir: None,
2070 };
2071
2072 let client = SocketClient::builder()
2073 .config(config)
2074 .connect()
2075 .await
2076 .expect("Client connect failed unexpectedly");
2077
2078 assert!(client.is_active(), "Client should start as active");
2080
2081 wait_until_async(|| async { client.is_active() }, Duration::from_secs(10)).await;
2084
2085 client
2086 .send_bytes(b"TestReconnect".into())
2087 .await
2088 .expect("Send failed");
2089
2090 client.close().await;
2091 server_task.abort();
2092 }
2093
2094 #[rstest]
2095 #[tokio::test]
2096 async fn test_state_sink_reports_initial_loss_and_recovery() {
2097 let (port, listener) = bind_test_server().await;
2098 let server_task = task::spawn(async move {
2099 let (mut socket, _) = listener.accept().await.unwrap();
2100 sleep(Duration::from_millis(100)).await;
2101 socket.shutdown().await.unwrap();
2102
2103 let (socket, _) = listener.accept().await.unwrap();
2104 run_echo_server(socket).await;
2105 });
2106 let config = SocketConfig {
2107 url: format!("127.0.0.1:{port}"),
2108 mode: Mode::Plain,
2109 suffix: b"\r\n".to_vec(),
2110 message_handler: None,
2111 heartbeat: None,
2112 connect_timeout_ms: Some(1_000),
2113 reconnect_delay_initial_ms: Some(10),
2114 reconnect_backoff_factor: Some(1.0),
2115 reconnect_delay_max_ms: Some(10),
2116 reconnect_jitter_ms: Some(0),
2117 reconnect_max_attempts: Some(3),
2118 connection_max_retries: Some(1),
2119 heartbeat_timeout_secs: None,
2120 certs_dir: None,
2121 };
2122 let states = Arc::new(BlockingMutex::new(Vec::new()));
2123 let states_callback = Arc::clone(&states);
2124 let sink = SocketStateSink::new(move |state| {
2125 states_callback.lock().push(state);
2126 });
2127
2128 let client = SocketClient::builder()
2129 .config(config)
2130 .state_sink(sink)
2131 .connect()
2132 .await
2133 .unwrap();
2134
2135 assert_eq!(*states.lock(), vec![SocketState::Connected]);
2136
2137 wait_until_async(
2138 || {
2139 let states = Arc::clone(&states);
2140 async move { states.lock().len() == 3 }
2141 },
2142 Duration::from_secs(5),
2143 )
2144 .await;
2145 assert_eq!(
2146 *states.lock(),
2147 vec![
2148 SocketState::Connected,
2149 SocketState::Disconnected,
2150 SocketState::Connected,
2151 ]
2152 );
2153
2154 client.close().await;
2155 assert_eq!(states.lock().len(), 3);
2156 server_task.abort();
2157 }
2158
2159 #[rstest]
2160 #[tokio::test]
2161 async fn test_state_sink_ignores_initial_connection_failure() {
2162 let (port, listener) = bind_test_server().await;
2163 drop(listener);
2164 let config = SocketConfig {
2165 url: format!("127.0.0.1:{port}"),
2166 mode: Mode::Plain,
2167 suffix: b"\r\n".to_vec(),
2168 message_handler: None,
2169 heartbeat: None,
2170 connect_timeout_ms: Some(100),
2171 reconnect_delay_initial_ms: Some(1),
2172 reconnect_backoff_factor: Some(1.0),
2173 reconnect_delay_max_ms: Some(1),
2174 reconnect_jitter_ms: Some(0),
2175 reconnect_max_attempts: Some(1),
2176 connection_max_retries: Some(1),
2177 heartbeat_timeout_secs: None,
2178 certs_dir: None,
2179 };
2180 let states = Arc::new(BlockingMutex::new(Vec::new()));
2181 let states_callback = Arc::clone(&states);
2182 let sink = SocketStateSink::new(move |state| {
2183 states_callback.lock().push(state);
2184 });
2185
2186 let result = SocketClient::builder()
2187 .config(config)
2188 .state_sink(sink)
2189 .connect()
2190 .await;
2191
2192 assert!(result.is_err());
2193 assert_eq!(*states.lock(), Vec::new());
2194 }
2195}
2196
2197#[cfg(test)]
2198#[cfg(not(feature = "turmoil"))]
2199#[cfg(not(all(feature = "simulation", madsim)))] mod rust_tests {
2201 use std::{
2202 pin::Pin,
2203 sync::{
2204 Arc,
2205 atomic::{AtomicBool, AtomicUsize},
2206 },
2207 task::{Context, Poll, Waker},
2208 };
2209
2210 use nautilus_common::testing::wait_until_async;
2211 use parking_lot::{Condvar, Mutex};
2212 use rstest::rstest;
2213 use tokio::{
2214 io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, DuplexStream, ReadBuf},
2215 net::TcpListener,
2216 sync::oneshot,
2217 task::{self, yield_now},
2218 time::{Duration, sleep},
2219 };
2220
2221 use super::*;
2222 use crate::{SocketState, socket::SocketHeartbeat};
2223
2224 const TEST_TIMEOUT: Duration = Duration::from_secs(10);
2225
2226 struct CondvarReleaseGuard {
2227 release: Arc<(Mutex<bool>, Condvar)>,
2228 }
2229
2230 impl CondvarReleaseGuard {
2231 fn new(release: Arc<(Mutex<bool>, Condvar)>) -> Self {
2232 Self { release }
2233 }
2234
2235 fn release(&self) {
2236 let (lock, condvar) = self.release.as_ref();
2237 let mut released = lock.lock();
2238 *released = true;
2239 condvar.notify_all();
2240 }
2241 }
2242
2243 impl Drop for CondvarReleaseGuard {
2244 fn drop(&mut self) {
2245 self.release();
2246 }
2247 }
2248
2249 async fn recv_rendezvous<T: Send + 'static>(
2250 receiver: std::sync::mpsc::Receiver<T>,
2251 name: &'static str,
2252 ) -> T {
2253 let receive_task = tokio::task::spawn_blocking(move || receiver.recv_timeout(TEST_TIMEOUT));
2254
2255 match tokio::time::timeout(TEST_TIMEOUT * 2, receive_task).await {
2256 Ok(Ok(Ok(value))) => value,
2257 Ok(Ok(Err(e))) => {
2258 panic!("{name} did not arrive within the test timeout: {e}")
2259 }
2260 Ok(Err(e)) => panic!("{name} receive task failed: {e}"),
2261 Err(e) => panic!("{name} receive task did not finish: {e}"),
2262 }
2263 }
2264
2265 async fn await_task_termination(
2266 task: tokio::task::JoinHandle<()>,
2267 name: &'static str,
2268 cancellation_ok: bool,
2269 ) {
2270 match tokio::time::timeout(TEST_TIMEOUT, task).await {
2271 Ok(Ok(())) => {}
2272 Ok(Err(e)) if cancellation_ok && e.is_cancelled() => {}
2273 Ok(Err(e)) => panic!("{name} failed: {e}"),
2274 Err(e) => panic!("{name} did not terminate within the test timeout: {e}"),
2275 }
2276 }
2277
2278 fn reconnect_test_config(port: u16) -> SocketConfig {
2279 SocketConfig {
2280 url: format!("127.0.0.1:{port}"),
2281 mode: Mode::Plain,
2282 suffix: b"\r\n".to_vec(),
2283 message_handler: None,
2284 heartbeat: None,
2285 connect_timeout_ms: Some(1_000),
2286 reconnect_delay_initial_ms: None,
2287 reconnect_backoff_factor: None,
2288 reconnect_delay_max_ms: None,
2289 reconnect_jitter_ms: None,
2290 connection_max_retries: Some(1),
2291 reconnect_max_attempts: None,
2292 heartbeat_timeout_secs: None,
2293 certs_dir: None,
2294 }
2295 }
2296
2297 #[rstest]
2298 #[tokio::test]
2299 async fn test_reconnect_outcome_is_aborted_before_connect() {
2300 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2301 let port = listener.local_addr().unwrap().port();
2302 let server = tokio::spawn(async move {
2303 let (_socket, _) = listener.accept().await.unwrap();
2304 std::future::pending::<()>().await;
2305 });
2306 let mut inner = SocketClientInner::connect_url(reconnect_test_config(port), None)
2307 .await
2308 .unwrap();
2309 inner
2310 .connection_mode
2311 .store(ConnectionMode::Disconnect.as_u8(), Ordering::SeqCst);
2312
2313 let outcome = inner.reconnect(None).await.unwrap();
2314
2315 assert_eq!(outcome, ReconnectOutcome::Aborted);
2316 server.abort();
2317 }
2318
2319 #[rstest]
2320 #[tokio::test]
2321 async fn test_reconnect_outcome_is_reconnected() {
2322 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2323 let port = listener.local_addr().unwrap().port();
2324 let server = tokio::spawn(async move {
2325 let (_first, _) = listener.accept().await.unwrap();
2326 let (_second, _) = listener.accept().await.unwrap();
2327 std::future::pending::<()>().await;
2328 });
2329 let mut inner = SocketClientInner::connect_url(reconnect_test_config(port), None)
2330 .await
2331 .unwrap();
2332 inner
2333 .connection_mode
2334 .store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
2335
2336 let outcome = inner.reconnect(None).await.unwrap();
2337
2338 assert_eq!(outcome, ReconnectOutcome::Reconnected);
2339 assert_eq!(
2340 ConnectionMode::from_atomic(&inner.connection_mode),
2341 ConnectionMode::Active
2342 );
2343 server.abort();
2344 }
2345
2346 #[rstest]
2347 #[tokio::test]
2348 async fn test_reconnect_aborts_retired_reader() {
2349 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2350 let port = listener.local_addr().unwrap().port();
2351 let mut inner = SocketClientInner::connect_url(reconnect_test_config(port), None)
2352 .await
2353 .unwrap();
2354 let (_peer, _) = listener.accept().await.unwrap();
2355 let (dropped_tx, dropped_rx) = tokio::sync::oneshot::channel::<()>();
2356
2357 let pending_reader = tokio::spawn(async move {
2358 let _sender = dropped_tx;
2359 std::future::pending::<()>().await;
2360 });
2361
2362 let previous_reader = std::mem::replace(&mut inner.read_task, pending_reader);
2363 previous_reader.abort();
2364 let _ = previous_reader.await;
2365 let old_fence = inner.read_fence.clone();
2366 inner
2367 .connection_mode
2368 .store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
2369
2370 let outcome = inner.reconnect(None).await.unwrap();
2371 let reader_result = tokio::time::timeout(Duration::from_secs(1), dropped_rx)
2372 .await
2373 .expect("retired reader task must be canceled");
2374
2375 assert_eq!(outcome, ReconnectOutcome::Reconnected);
2376 assert!(
2377 reader_result.is_err(),
2378 "retired reader must drop its sender without sending"
2379 );
2380 assert!(!old_fence.is_valid());
2381 assert!(inner.read_fence.is_valid());
2382 assert_eq!(
2383 ConnectionMode::from_atomic(&inner.connection_mode),
2384 ConnectionMode::Active
2385 );
2386 }
2387
2388 #[rstest]
2389 #[tokio::test]
2390 async fn test_inner_drop_invalidates_read_fence_and_aborts_tasks() {
2391 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2392 let port = listener.local_addr().unwrap().port();
2393 let server = tokio::spawn(async move {
2394 let (_socket, _) = listener.accept().await.unwrap();
2395 std::future::pending::<()>().await;
2396 });
2397 let mut config = reconnect_test_config(port);
2398 config.heartbeat = Some(SocketHeartbeat {
2399 interval_secs: 60,
2400 payload: b"ping\r\n".to_vec(),
2401 });
2402 let inner = SocketClientInner::connect_url(config, None).await.unwrap();
2403 let read_fence = inner.read_fence.clone();
2404 let read_abort = inner.read_task.abort_handle();
2405 let write_abort = inner.write_task.abort_handle();
2406 let heartbeat_abort = inner
2407 .heartbeat_task
2408 .as_ref()
2409 .expect("heartbeat task should be spawned for a configured heartbeat")
2410 .abort_handle();
2411
2412 assert!(read_fence.is_valid(), "read fence should start valid");
2413 assert!(
2414 !read_abort.is_finished(),
2415 "read task should be running before drop"
2416 );
2417 assert!(
2418 !write_abort.is_finished(),
2419 "write task should be running before drop"
2420 );
2421 assert!(
2422 !heartbeat_abort.is_finished(),
2423 "heartbeat task should be running before drop"
2424 );
2425
2426 drop(inner);
2427 wait_until_async(
2428 || async {
2429 read_abort.is_finished()
2430 && write_abort.is_finished()
2431 && heartbeat_abort.is_finished()
2432 },
2433 TEST_TIMEOUT,
2434 )
2435 .await;
2436
2437 assert!(!read_fence.is_valid(), "read fence was not invalidated");
2438 assert!(read_abort.is_finished(), "read task was not aborted");
2439 assert!(write_abort.is_finished(), "write task was not aborted");
2440 assert!(
2441 heartbeat_abort.is_finished(),
2442 "heartbeat task was not aborted"
2443 );
2444 server.abort();
2445 }
2446
2447 struct ScriptedReader {
2448 first: Option<Vec<u8>>,
2449 remainder: Arc<Mutex<Option<Vec<u8>>>>,
2450 pending_tx: Option<oneshot::Sender<()>>,
2451 waker: Arc<Mutex<Option<Waker>>>,
2452 }
2453
2454 impl AsyncRead for ScriptedReader {
2455 fn poll_read(
2456 mut self: Pin<&mut Self>,
2457 cx: &mut Context<'_>,
2458 buf: &mut ReadBuf<'_>,
2459 ) -> Poll<std::io::Result<()>> {
2460 if let Some(first) = self.first.take() {
2461 buf.put_slice(&first);
2462 return Poll::Ready(Ok(()));
2463 }
2464
2465 if let Some(remainder) = self.remainder.lock().take() {
2466 buf.put_slice(&remainder);
2467 return Poll::Ready(Ok(()));
2468 }
2469
2470 if let Some(tx) = self.pending_tx.take() {
2471 let _ = tx.send(());
2472 }
2473 *self.waker.lock() = Some(cx.waker().clone());
2474 Poll::Pending
2475 }
2476 }
2477
2478 struct LimitProbeReader {
2479 remaining: usize,
2480 read_after_limit: Arc<AtomicBool>,
2481 }
2482
2483 impl AsyncRead for LimitProbeReader {
2484 fn poll_read(
2485 mut self: Pin<&mut Self>,
2486 _cx: &mut Context<'_>,
2487 buf: &mut ReadBuf<'_>,
2488 ) -> Poll<std::io::Result<()>> {
2489 if self.remaining == 0 {
2490 self.read_after_limit.store(true, Ordering::SeqCst);
2491 return Poll::Ready(Ok(()));
2492 }
2493
2494 let read_len = self.remaining.min(buf.remaining());
2495 buf.initialize_unfilled_to(read_len).fill(b'x');
2496 buf.advance(read_len);
2497 self.remaining -= read_len;
2498 Poll::Ready(Ok(()))
2499 }
2500 }
2501
2502 struct BackpressuredWriter {
2503 stream: DuplexStream,
2504 pending_tx: Option<oneshot::Sender<()>>,
2505 }
2506
2507 impl AsyncWrite for BackpressuredWriter {
2508 fn poll_write(
2509 mut self: Pin<&mut Self>,
2510 cx: &mut Context<'_>,
2511 buf: &[u8],
2512 ) -> Poll<std::io::Result<usize>> {
2513 let result = Pin::new(&mut self.stream).poll_write(cx, buf);
2514 if result.is_pending()
2515 && let Some(tx) = self.pending_tx.take()
2516 {
2517 let _ = tx.send(());
2518 }
2519 result
2520 }
2521
2522 fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
2523 Pin::new(&mut self.stream).poll_flush(cx)
2524 }
2525
2526 fn poll_shutdown(
2527 mut self: Pin<&mut Self>,
2528 cx: &mut Context<'_>,
2529 ) -> Poll<std::io::Result<()>> {
2530 Pin::new(&mut self.stream).poll_shutdown(cx)
2531 }
2532 }
2533
2534 struct FailingWriter;
2535
2536 impl AsyncWrite for FailingWriter {
2537 fn poll_write(
2538 self: Pin<&mut Self>,
2539 _cx: &mut Context<'_>,
2540 _buf: &[u8],
2541 ) -> Poll<std::io::Result<usize>> {
2542 Poll::Ready(Err(std::io::Error::new(
2543 std::io::ErrorKind::BrokenPipe,
2544 "test writer failure",
2545 )))
2546 }
2547
2548 fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
2549 Poll::Ready(Ok(()))
2550 }
2551
2552 fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
2553 Poll::Ready(Ok(()))
2554 }
2555 }
2556
2557 struct RecordingWriter {
2558 bytes: Arc<Mutex<Vec<u8>>>,
2559 }
2560
2561 impl AsyncWrite for RecordingWriter {
2562 fn poll_write(
2563 self: Pin<&mut Self>,
2564 _cx: &mut Context<'_>,
2565 buf: &[u8],
2566 ) -> Poll<std::io::Result<usize>> {
2567 self.bytes.lock().extend_from_slice(buf);
2568 Poll::Ready(Ok(buf.len()))
2569 }
2570
2571 fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
2572 Poll::Ready(Ok(()))
2573 }
2574
2575 fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
2576 Poll::Ready(Ok(()))
2577 }
2578 }
2579
2580 fn test_socket_client(
2581 connection_state: Arc<AtomicU8>,
2582 state_notify: Arc<tokio::sync::Notify>,
2583 controller_task: tokio::task::JoinHandle<()>,
2584 ) -> SocketClient {
2585 let (writer_tx, _writer_rx) = tokio::sync::mpsc::unbounded_channel();
2586 let controller_lifecycle = Arc::new(ControllerLifecycle::new());
2587 let abort_handle = controller_task.abort_handle();
2588 controller_lifecycle.set_abort(move || abort_handle.abort());
2589
2590 SocketClient {
2591 controller_task,
2592 connection_mode: connection_state,
2593 state_notify,
2594 connect_timeout: Duration::from_secs(1),
2595 writer_tx,
2596 controller_lifecycle,
2597 controller_notify: Arc::new(tokio::sync::Notify::new()),
2598 state_sink: None,
2599 }
2600 }
2601
2602 #[rstest]
2603 #[tokio::test]
2604 async fn test_reconnect_handle_is_closed_after_client_drop() {
2605 let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
2606 let state_notify = Arc::new(tokio::sync::Notify::new());
2607 let controller_task = tokio::spawn(std::future::pending::<()>());
2608 let client = test_socket_client(connection_state, state_notify, controller_task);
2609 let handle = client.reconnect_handle();
2610
2611 drop(client);
2612
2613 assert_eq!(handle.request_reconnect(), ReconnectRequestOutcome::Closed);
2614 }
2615
2616 #[rstest]
2617 #[tokio::test]
2618 async fn test_concurrent_drop_defers_controller_abort_until_request_completes() {
2619 let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
2620 let state_notify = Arc::new(tokio::sync::Notify::new());
2621
2622 let controller_task = tokio::spawn(std::future::pending::<()>());
2623 let controller_abort = controller_task.abort_handle();
2624 let mut client = test_socket_client(connection_state, state_notify, controller_task);
2625 let callback_release = Arc::new((Mutex::new(false), Condvar::new()));
2626 let callback_release_guard = CondvarReleaseGuard::new(Arc::clone(&callback_release));
2627 let callback_release_clone = Arc::clone(&callback_release);
2628 let (callback_entered_tx, callback_entered_rx) = std::sync::mpsc::channel();
2629 let states = Arc::new(Mutex::new(Vec::new()));
2630 let states_clone = Arc::clone(&states);
2631 client.state_sink = Some(SocketStateSink::new(move |state| {
2632 states_clone.lock().push(state);
2633 callback_entered_tx.send(()).unwrap();
2634 let (lock, condvar) = callback_release_clone.as_ref();
2635 let mut released = lock.lock();
2636
2637 while !*released {
2638 condvar.wait(&mut released);
2639 }
2640 }));
2641 let handle = client.reconnect_handle();
2642 let surviving_handle = handle.clone();
2643 let controller_notify = Arc::clone(&handle.controller_notify);
2644 let request_thread = std::thread::spawn(move || handle.request_reconnect());
2645
2646 recv_rendezvous(callback_entered_rx, "socket state callback entry").await;
2647 let (drop_finished_tx, drop_finished_rx) = std::sync::mpsc::channel();
2648 let drop_thread = std::thread::spawn(move || {
2649 drop(client);
2650 drop_finished_tx.send(()).unwrap();
2651 });
2652
2653 recv_rendezvous(drop_finished_rx, "client drop").await;
2654 drop_thread.join().unwrap();
2655 assert!(!controller_abort.is_finished());
2656
2657 callback_release_guard.release();
2658 assert_eq!(
2659 request_thread.join().unwrap(),
2660 ReconnectRequestOutcome::Accepted
2661 );
2662 tokio::time::timeout(Duration::from_millis(10), controller_notify.notified())
2663 .await
2664 .expect("accepted request should notify before deferred controller abort");
2665 wait_until_async(|| async { controller_abort.is_finished() }, TEST_TIMEOUT).await;
2666
2667 assert_eq!(
2668 surviving_handle.request_reconnect(),
2669 ReconnectRequestOutcome::Closed
2670 );
2671 assert_eq!(*states.lock(), vec![SocketState::Disconnected]);
2672 assert!(
2673 tokio::time::timeout(Duration::from_millis(10), controller_notify.notified())
2674 .await
2675 .is_err(),
2676 "closed request should not notify controller",
2677 );
2678 }
2679
2680 #[rstest]
2681 #[tokio::test]
2682 async fn test_reconnect_state_callback_can_drop_client() {
2683 let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
2684 let state_notify = Arc::new(tokio::sync::Notify::new());
2685
2686 let controller_task = tokio::spawn(std::future::pending::<()>());
2687 let controller_abort = controller_task.abort_handle();
2688 let mut client = test_socket_client(connection_state, state_notify, controller_task);
2689 let client_slot = Arc::new(Mutex::new(None));
2690 let client_slot_callback = Arc::clone(&client_slot);
2691 let states = Arc::new(Mutex::new(Vec::new()));
2692 let states_callback = Arc::clone(&states);
2693 client.state_sink = Some(SocketStateSink::new(move |state| {
2694 states_callback.lock().push(state);
2695 drop(client_slot_callback.lock().take());
2696 }));
2697 let handle = client.reconnect_handle();
2698 let controller_notify = Arc::clone(&handle.controller_notify);
2699 *client_slot.lock() = Some(client);
2700 let (result_tx, result_rx) = std::sync::mpsc::channel();
2701 std::thread::spawn(move || result_tx.send(handle.request_reconnect()).unwrap());
2702
2703 assert_eq!(
2704 recv_rendezvous(result_rx, "reconnect request").await,
2705 ReconnectRequestOutcome::Accepted
2706 );
2707 tokio::time::timeout(Duration::from_millis(10), controller_notify.notified())
2708 .await
2709 .expect("accepted request should notify before deferred controller abort");
2710 wait_until_async(|| async { controller_abort.is_finished() }, TEST_TIMEOUT).await;
2711
2712 assert!(client_slot.lock().is_none());
2713 assert_eq!(*states.lock(), vec![SocketState::Disconnected]);
2714 }
2715
2716 #[rstest]
2717 #[tokio::test]
2718 async fn test_manual_reconnect_uses_controller_path() {
2719 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2720 let port = listener.local_addr().unwrap().port();
2721 let (first_accepted_tx, first_accepted_rx) = oneshot::channel();
2722 let (second_accepted_tx, second_accepted_rx) = oneshot::channel();
2723 let (payload_tx, payload_rx) = oneshot::channel();
2724 let (release_tx, release_rx) = oneshot::channel();
2725
2726 let server = task::spawn(async move {
2727 let (_first, _) = listener.accept().await.unwrap();
2728 first_accepted_tx.send(()).unwrap();
2729
2730 let (mut second, _) = listener.accept().await.unwrap();
2731 second_accepted_tx.send(()).unwrap();
2732 let mut payload = [0_u8; 8];
2733 second.read_exact(&mut payload).await.unwrap();
2734 payload_tx.send(payload).unwrap();
2735 let _ = release_rx.await;
2736 });
2737 let states = Arc::new(Mutex::new(Vec::new()));
2738 let states_callback = Arc::clone(&states);
2739 let sink = SocketStateSink::new(move |state| {
2740 states_callback.lock().push(state);
2741 });
2742 let callback_count = Arc::new(AtomicUsize::new(0));
2743 let callback_count_clone = Arc::clone(&callback_count);
2744 let post_reconnection = Arc::new(move || {
2745 callback_count_clone.fetch_add(1, Ordering::SeqCst);
2746 });
2747 let client = SocketClient::builder()
2748 .config(reconnect_test_config(port))
2749 .post_reconnection(post_reconnection)
2750 .state_sink(sink)
2751 .connect()
2752 .await
2753 .unwrap();
2754 first_accepted_rx.await.unwrap();
2755 let handle = client.reconnect_handle();
2756
2757 assert!(client.request_reconnect());
2758 assert_eq!(
2759 handle.request_reconnect(),
2760 ReconnectRequestOutcome::AlreadyReconnecting
2761 );
2762 assert!(!client.request_reconnect());
2763 assert_eq!(
2764 *states.lock(),
2765 vec![SocketState::Connected, SocketState::Disconnected]
2766 );
2767
2768 tokio::time::timeout(TEST_TIMEOUT, second_accepted_rx)
2769 .await
2770 .expect("controller should establish a replacement connection")
2771 .unwrap();
2772 wait_until_async(
2773 || async { client.is_active() && callback_count.load(Ordering::SeqCst) == 1 },
2774 TEST_TIMEOUT,
2775 )
2776 .await;
2777 client.send_bytes(b"manual".to_vec()).await.unwrap();
2778
2779 assert_eq!(
2780 tokio::time::timeout(TEST_TIMEOUT, payload_rx)
2781 .await
2782 .expect("replacement connection should receive the framed payload")
2783 .unwrap(),
2784 *b"manual\r\n"
2785 );
2786 assert_eq!(callback_count.load(Ordering::SeqCst), 1);
2787 assert_eq!(
2788 *states.lock(),
2789 vec![
2790 SocketState::Connected,
2791 SocketState::Disconnected,
2792 SocketState::Connected,
2793 ]
2794 );
2795
2796 client.close().await;
2797 release_tx.send(()).unwrap();
2798 server.await.unwrap();
2799 assert_eq!(callback_count.load(Ordering::SeqCst), 1);
2800 assert_eq!(states.lock().len(), 3);
2801 }
2802
2803 #[rstest]
2804 #[case(ConnectionMode::Disconnect)]
2805 #[case(ConnectionMode::Closed)]
2806 #[tokio::test]
2807 async fn test_send_bytes_rejects_terminal_state(#[case] mode: ConnectionMode) {
2808 let connection_state = Arc::new(AtomicU8::new(mode.as_u8()));
2809 let state_notify = Arc::new(tokio::sync::Notify::new());
2810 let controller_task = tokio::spawn(std::future::pending::<()>());
2811 let client = test_socket_client(connection_state, state_notify, controller_task);
2812
2813 let result = client.send_bytes(b"terminal".to_vec()).await;
2814
2815 assert!(matches!(result, Err(SendError::Closed)));
2816 }
2817
2818 #[rstest]
2819 #[tokio::test(start_paused = true)]
2820 async fn test_close_sets_closed_after_controller_was_aborted() {
2821 let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
2822 let state_notify = Arc::new(tokio::sync::Notify::new());
2823
2824 let controller_task = tokio::spawn(std::future::pending::<()>());
2825 controller_task.abort();
2826 let client = test_socket_client(connection_state, state_notify, controller_task);
2827
2828 client.close_with_timeout(Duration::from_millis(1)).await;
2829
2830 assert!(client.is_closed());
2831 }
2832
2833 #[rstest]
2834 #[tokio::test]
2835 async fn test_reconnect_exhaustion_closes_client() {
2836 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2837 let port = listener.local_addr().unwrap().port();
2838 let (accepted_tx, accepted_rx) = oneshot::channel();
2839 let (release_tx, release_rx) = oneshot::channel();
2840
2841 let server = task::spawn(async move {
2842 let (socket, _) = listener.accept().await.unwrap();
2843 let _ = accepted_tx.send(());
2844 let _ = release_rx.await;
2845 drop(socket);
2846 drop(listener);
2847 });
2848
2849 let config = SocketConfig {
2850 url: format!("127.0.0.1:{port}"),
2851 mode: Mode::Plain,
2852 suffix: b"\r\n".to_vec(),
2853 message_handler: None,
2854 heartbeat: None,
2855 connect_timeout_ms: Some(100),
2856 reconnect_delay_initial_ms: Some(1),
2857 reconnect_delay_max_ms: Some(1),
2858 reconnect_backoff_factor: Some(1.0),
2859 reconnect_jitter_ms: Some(0),
2860 connection_max_retries: Some(1),
2861 reconnect_max_attempts: Some(1),
2862 heartbeat_timeout_secs: None,
2863 certs_dir: None,
2864 };
2865
2866 let states = Arc::new(Mutex::new(Vec::new()));
2867 let states_callback = Arc::clone(&states);
2868 let sink = SocketStateSink::new(move |state| {
2869 states_callback.lock().push(state);
2870 });
2871 let client = SocketClient::builder()
2872 .config(config)
2873 .state_sink(sink)
2874 .connect()
2875 .await
2876 .unwrap();
2877 accepted_rx.await.unwrap();
2878 release_tx.send(()).unwrap();
2879
2880 wait_until_async(|| async { client.is_closed() }, Duration::from_secs(5)).await;
2881 assert!(client.is_closed());
2882 assert_eq!(
2883 *states.lock(),
2884 vec![SocketState::Connected, SocketState::Disconnected]
2885 );
2886
2887 client.close().await;
2888 assert!(client.is_closed());
2889 assert_eq!(states.lock().len(), 2);
2890 server.await.unwrap();
2891 }
2892
2893 #[rstest]
2894 #[tokio::test]
2895 async fn test_drop_suppresses_socket_state_event() {
2896 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2897 let port = listener.local_addr().unwrap().port();
2898 let (accepted_tx, accepted_rx) = oneshot::channel();
2899
2900 let server = task::spawn(async move {
2901 let (_socket, _) = listener.accept().await.unwrap();
2902 let _ = accepted_tx.send(());
2903 std::future::pending::<()>().await;
2904 });
2905
2906 let config = SocketConfig {
2907 url: format!("127.0.0.1:{port}"),
2908 mode: Mode::Plain,
2909 suffix: b"\r\n".to_vec(),
2910 message_handler: None,
2911 heartbeat: None,
2912 connect_timeout_ms: Some(100),
2913 reconnect_delay_initial_ms: Some(1),
2914 reconnect_delay_max_ms: Some(1),
2915 reconnect_backoff_factor: Some(1.0),
2916 reconnect_jitter_ms: Some(0),
2917 connection_max_retries: Some(1),
2918 reconnect_max_attempts: Some(1),
2919 heartbeat_timeout_secs: None,
2920 certs_dir: None,
2921 };
2922 let states = Arc::new(Mutex::new(Vec::new()));
2923 let states_callback = Arc::clone(&states);
2924 let sink = SocketStateSink::new(move |state| {
2925 states_callback.lock().push(state);
2926 });
2927 let client = SocketClient::builder()
2928 .config(config)
2929 .state_sink(sink)
2930 .connect()
2931 .await
2932 .unwrap();
2933 accepted_rx.await.unwrap();
2934
2935 drop(client);
2936 sleep(Duration::from_millis(25)).await;
2937
2938 assert_eq!(*states.lock(), vec![SocketState::Connected]);
2939 server.abort();
2940 }
2941
2942 #[rstest]
2943 #[tokio::test]
2944 async fn test_graceful_close_is_idempotent() {
2945 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2946 let port = listener.local_addr().unwrap().port();
2947 let (accepted_tx, accepted_rx) = oneshot::channel();
2948
2949 let server = task::spawn(async move {
2950 let (_socket, _) = listener.accept().await.unwrap();
2951 let _ = accepted_tx.send(());
2952 std::future::pending::<()>().await;
2953 });
2954 let config = SocketConfig {
2955 url: format!("127.0.0.1:{port}"),
2956 mode: Mode::Plain,
2957 suffix: b"\r\n".to_vec(),
2958 message_handler: None,
2959 heartbeat: None,
2960 connect_timeout_ms: None,
2961 reconnect_delay_initial_ms: None,
2962 reconnect_delay_max_ms: None,
2963 reconnect_backoff_factor: None,
2964 reconnect_jitter_ms: None,
2965 connection_max_retries: None,
2966 reconnect_max_attempts: None,
2967 heartbeat_timeout_secs: None,
2968 certs_dir: None,
2969 };
2970 let client = SocketClient::builder()
2971 .config(config)
2972 .connect()
2973 .await
2974 .unwrap();
2975 accepted_rx.await.unwrap();
2976
2977 client.close().await;
2978 client.close().await;
2979
2980 assert!(client.is_closed());
2981 server.abort();
2982 }
2983
2984 #[rstest]
2985 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
2986 async fn test_read_loop_drops_remaining_frames_after_session_replaced() {
2987 let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
2988 let read_fence = ReadSessionFence::new();
2989 let received = Arc::new(Mutex::new(Vec::<Vec<u8>>::new()));
2990 let received_clone = Arc::clone(&received);
2991 let handler_release = Arc::new((Mutex::new(false), Condvar::new()));
2992 let handler_release_guard = CondvarReleaseGuard::new(Arc::clone(&handler_release));
2993 let handler_release_clone = Arc::clone(&handler_release);
2994 let (handler_entered_tx, handler_entered_rx) = std::sync::mpsc::channel();
2995 let handler: TcpMessageHandler = Arc::new(move |data| {
2996 received_clone.lock().push(data.to_vec());
2997 handler_entered_tx.send(()).unwrap();
2998 let (lock, condvar) = handler_release_clone.as_ref();
2999 let mut released = lock.lock();
3000
3001 while !*released {
3002 condvar.wait(&mut released);
3003 }
3004 });
3005 let reader = ScriptedReader {
3006 first: Some(b"first\r\nsecond\r\n".to_vec()),
3007 remainder: Arc::new(Mutex::new(None)),
3008 pending_tx: None,
3009 waker: Arc::new(Mutex::new(None)),
3010 };
3011 let read_task = SocketClientInner::spawn_read_task(
3012 Arc::clone(&connection_state),
3013 read_fence.clone(),
3014 reader,
3015 Some(handler),
3016 b"\r\n".to_vec(),
3017 None,
3018 );
3019
3020 recv_rendezvous(handler_entered_rx, "socket first-handler entry").await;
3021 connection_state.store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
3022 read_fence.invalidate();
3023 read_task.abort();
3024 connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
3025
3026 handler_release_guard.release();
3027 await_task_termination(read_task, "old socket read task", true).await;
3028
3029 assert_eq!(received.lock().as_slice(), &[b"first".to_vec()]);
3030 }
3031
3032 #[rstest]
3033 #[tokio::test(start_paused = true)]
3034 async fn test_read_loop_drops_partial_old_session_frame() {
3035 let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
3036 let old_read_fence = ReadSessionFence::new();
3037 let received = Arc::new(Mutex::new(Vec::<Vec<u8>>::new()));
3038 let received_clone = Arc::clone(&received);
3039 let handler: TcpMessageHandler =
3040 Arc::new(move |data| received_clone.lock().push(data.to_vec()));
3041 let (pending_tx, pending_rx) = oneshot::channel();
3042 let remainder = Arc::new(Mutex::new(None));
3043 let waker = Arc::new(Mutex::new(None));
3044 let reader = ScriptedReader {
3045 first: Some(b"old".to_vec()),
3046 remainder: Arc::clone(&remainder),
3047 pending_tx: Some(pending_tx),
3048 waker: Arc::clone(&waker),
3049 };
3050
3051 let read_task = SocketClientInner::spawn_read_task(
3052 Arc::clone(&connection_state),
3053 old_read_fence.clone(),
3054 reader,
3055 Some(handler),
3056 b"\r\n".to_vec(),
3057 None,
3058 );
3059
3060 pending_rx.await.unwrap();
3061 connection_state.store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
3062 old_read_fence.invalidate();
3063 connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
3064 *remainder.lock() = Some(b"\r\nnew\r\n".to_vec());
3065 waker.lock().take().unwrap().wake();
3066 read_task.await.unwrap();
3067
3068 let fresh_reader = ScriptedReader {
3069 first: Some(b"new\r\n".to_vec()),
3070 remainder: Arc::new(Mutex::new(None)),
3071 pending_tx: None,
3072 waker: Arc::new(Mutex::new(None)),
3073 };
3074 let fresh_connection_state = Arc::clone(&connection_state);
3075 let fresh_received = Arc::clone(&received);
3076 let fresh_handler: TcpMessageHandler = Arc::new(move |data| {
3077 fresh_received.lock().push(data.to_vec());
3078 fresh_connection_state.store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
3079 });
3080 SocketClientInner::spawn_read_task(
3081 Arc::clone(&connection_state),
3082 ReadSessionFence::new(),
3083 fresh_reader,
3084 Some(fresh_handler),
3085 b"\r\n".to_vec(),
3086 None,
3087 )
3088 .await
3089 .unwrap();
3090
3091 assert_eq!(received.lock().as_slice(), &[b"new".to_vec()]);
3092 }
3093
3094 #[rstest]
3095 #[tokio::test(start_paused = true)]
3096 async fn test_read_loop_stops_when_first_handler_ends_session() {
3097 let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
3098 let handler_state = Arc::clone(&connection_state);
3099 let received = Arc::new(Mutex::new(Vec::<Vec<u8>>::new()));
3100 let received_clone = Arc::clone(&received);
3101 let handler: TcpMessageHandler = Arc::new(move |data| {
3102 received_clone.lock().push(data.to_vec());
3103 handler_state.store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
3104 });
3105 let (pending_tx, _pending_rx) = oneshot::channel();
3106 let reader = ScriptedReader {
3107 first: Some(b"first\r\nsecond\r\n".to_vec()),
3108 remainder: Arc::new(Mutex::new(None)),
3109 pending_tx: Some(pending_tx),
3110 waker: Arc::new(Mutex::new(None)),
3111 };
3112
3113 SocketClientInner::spawn_read_task(
3114 connection_state,
3115 ReadSessionFence::new(),
3116 reader,
3117 Some(handler),
3118 b"\r\n".to_vec(),
3119 None,
3120 )
3121 .await
3122 .unwrap();
3123
3124 assert_eq!(received.lock().as_slice(), &[b"first".to_vec()]);
3125 }
3126
3127 #[tokio::test]
3128 async fn test_read_loop_closes_when_unframed_buffer_exceeds_limit() {
3129 let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
3130 let read_after_limit = Arc::new(AtomicBool::new(false));
3131 let reader = LimitProbeReader {
3132 remaining: MAX_READ_BUFFER_BYTES + 1,
3133 read_after_limit: Arc::clone(&read_after_limit),
3134 };
3135
3136 SocketClientInner::run_read_loop(
3137 connection_state,
3138 ReadSessionFence::new(),
3139 reader,
3140 None,
3141 b"\r\n".to_vec(),
3142 None,
3143 Duration::from_secs(1),
3144 )
3145 .await;
3146
3147 assert!(!read_after_limit.load(Ordering::SeqCst));
3148 }
3149
3150 #[rstest]
3151 #[tokio::test(start_paused = true)]
3152 async fn test_stalled_socket_write_sends_reconnect_replay_before_buffer() {
3153 type TestWriter = Pin<Box<dyn AsyncWrite + Send>>;
3154
3155 let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
3156 let state_notify = Arc::new(tokio::sync::Notify::new());
3157 let (stream, _non_reading_peer) = tokio::io::duplex(1);
3158 let (pending_tx, pending_rx) = oneshot::channel();
3159 let writer: TestWriter = Box::pin(BackpressuredWriter {
3160 stream,
3161 pending_tx: Some(pending_tx),
3162 });
3163 let (writer_tx, writer_rx) =
3164 tokio::sync::mpsc::unbounded_channel::<WriterCommand<TestWriter>>();
3165 let states = Arc::new(Mutex::new(Vec::new()));
3166 let states_callback = Arc::clone(&states);
3167 let sink = SocketStateSink::new(move |state| {
3168 states_callback.lock().push(state);
3169 });
3170 let write_task = SocketClientInner::spawn_write_task(
3171 Arc::clone(&connection_state),
3172 Arc::clone(&state_notify),
3173 writer,
3174 writer_rx,
3175 b"\r\n".to_vec(),
3176 Some(sink),
3177 );
3178
3179 writer_tx
3180 .send(WriterCommand::Send(Bytes::from_static(b"complete-message")))
3181 .unwrap();
3182 pending_rx.await.unwrap();
3183
3184 let recorded = Arc::new(Mutex::new(Vec::new()));
3185 let new_writer: TestWriter = Box::pin(RecordingWriter {
3186 bytes: Arc::clone(&recorded),
3187 });
3188 let (update_tx, update_rx) = oneshot::channel();
3189 writer_tx
3190 .send(WriterCommand::UpdateWithReplay(
3191 new_writer,
3192 vec![Bytes::from_static(b"authentication")],
3193 update_tx,
3194 ))
3195 .unwrap();
3196
3197 tokio::time::advance(Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS)).await;
3198 yield_now().await;
3199
3200 assert!(
3201 tokio::time::timeout(Duration::from_secs(10), update_rx)
3202 .await
3203 .expect("writer update should not remain queued behind a stalled write")
3204 .unwrap(),
3205 "writer should report successful replay"
3206 );
3207 assert_eq!(
3208 ConnectionMode::from_atomic(&connection_state),
3209 ConnectionMode::Reconnect
3210 );
3211 assert_eq!(
3212 recorded.lock().as_slice(),
3213 b"authentication\r\ncomplete-message\r\n"
3214 );
3215
3216 connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
3217 state_notify.notify_waiters();
3218 drop(writer_tx);
3219 write_task.await.unwrap();
3220
3221 assert_eq!(*states.lock(), vec![SocketState::Disconnected]);
3222 }
3223
3224 #[rstest]
3225 #[tokio::test(start_paused = true)]
3226 async fn test_send_after_writer_update_drains_when_reconnect_completes() {
3227 type TestWriter = Pin<Box<dyn AsyncWrite + Send>>;
3228
3229 let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Reconnect.as_u8()));
3230 let state_notify = Arc::new(tokio::sync::Notify::new());
3231 let initial_writer: TestWriter = Box::pin(RecordingWriter {
3232 bytes: Arc::new(Mutex::new(Vec::new())),
3233 });
3234 let (writer_tx, writer_rx) =
3235 tokio::sync::mpsc::unbounded_channel::<WriterCommand<TestWriter>>();
3236 let write_task = SocketClientInner::spawn_write_task(
3237 Arc::clone(&connection_state),
3238 Arc::clone(&state_notify),
3239 initial_writer,
3240 writer_rx,
3241 b"\r\n".to_vec(),
3242 None,
3243 );
3244
3245 let recorded = Arc::new(Mutex::new(Vec::new()));
3246 let new_writer: TestWriter = Box::pin(RecordingWriter {
3247 bytes: Arc::clone(&recorded),
3248 });
3249 let (update_tx, update_rx) = oneshot::channel();
3250 writer_tx
3251 .send(WriterCommand::Update(new_writer, update_tx))
3252 .unwrap();
3253 assert!(update_rx.await.unwrap());
3254
3255 writer_tx
3256 .send(WriterCommand::Send(Bytes::from_static(b"late")))
3257 .unwrap();
3258 yield_now().await;
3259
3260 connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
3261 state_notify.notify_waiters();
3262 writer_tx
3263 .send(WriterCommand::Send(Bytes::from_static(b"new")))
3264 .unwrap();
3265 tokio::time::advance(Duration::from_millis(
3266 CONNECTION_STATE_CHECK_INTERVAL_MS * 2,
3267 ))
3268 .await;
3269 yield_now().await;
3270
3271 let actual = recorded.lock().clone();
3272 connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
3273 state_notify.notify_waiters();
3274 drop(writer_tx);
3275 write_task.await.unwrap();
3276
3277 assert_eq!(actual, b"late\r\nnew\r\n");
3278 }
3279
3280 #[rstest]
3281 #[case::latest_wins(
3282 &["subscription-a", "subscription-b"],
3283 &["subscription-b"],
3284 b"subscription-b\r\n",
3285 )]
3286 #[case::different_replay_drains(
3287 &["subscription-b"],
3288 &["subscription-a"],
3289 b"subscription-a\r\nsubscription-b\r\n",
3290 )]
3291 #[tokio::test]
3292 async fn test_send_or_replay_buffering(
3293 #[case] buffered: &[&str],
3294 #[case] replay: &[&str],
3295 #[case] expected: &[u8],
3296 ) {
3297 type TestWriter = Pin<Box<dyn AsyncWrite + Send>>;
3298
3299 let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Reconnect.as_u8()));
3300 let state_notify = Arc::new(tokio::sync::Notify::new());
3301 let initial_writer: TestWriter = Box::pin(RecordingWriter {
3302 bytes: Arc::new(Mutex::new(Vec::new())),
3303 });
3304 let (writer_tx, writer_rx) =
3305 tokio::sync::mpsc::unbounded_channel::<WriterCommand<TestWriter>>();
3306 let write_task = SocketClientInner::spawn_write_task(
3307 Arc::clone(&connection_state),
3308 Arc::clone(&state_notify),
3309 initial_writer,
3310 writer_rx,
3311 b"\r\n".to_vec(),
3312 None,
3313 );
3314
3315 for data in buffered {
3316 writer_tx
3317 .send(WriterCommand::SendOrReplay {
3318 key: 7,
3319 data: Bytes::copy_from_slice(data.as_bytes()),
3320 })
3321 .unwrap();
3322 }
3323 let recorded = Arc::new(Mutex::new(Vec::new()));
3324 let new_writer: TestWriter = Box::pin(RecordingWriter {
3325 bytes: Arc::clone(&recorded),
3326 });
3327 let (update_tx, update_rx) = oneshot::channel();
3328 writer_tx
3329 .send(WriterCommand::UpdateWithReplay(
3330 new_writer,
3331 replay
3332 .iter()
3333 .map(|data| Bytes::copy_from_slice(data.as_bytes()))
3334 .collect(),
3335 update_tx,
3336 ))
3337 .unwrap();
3338
3339 assert!(update_rx.await.unwrap());
3340 assert_eq!(recorded.lock().as_slice(), expected);
3341 connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
3342 state_notify.notify_waiters();
3343 drop(writer_tx);
3344 write_task.await.unwrap();
3345 }
3346
3347 #[rstest]
3348 #[tokio::test(start_paused = true)]
3349 async fn test_active_reconnect_buffer_write_failure_reconnects_and_retries() {
3350 type TestWriter = Pin<Box<dyn AsyncWrite + Send>>;
3351
3352 let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Reconnect.as_u8()));
3353 let state_notify = Arc::new(tokio::sync::Notify::new());
3354 let initial_writer: TestWriter = Box::pin(RecordingWriter {
3355 bytes: Arc::new(Mutex::new(Vec::new())),
3356 });
3357 let (writer_tx, writer_rx) =
3358 tokio::sync::mpsc::unbounded_channel::<WriterCommand<TestWriter>>();
3359 let states = Arc::new(Mutex::new(Vec::new()));
3360 let states_callback = Arc::clone(&states);
3361 let sink = SocketStateSink::new(move |state| {
3362 states_callback.lock().push(state);
3363 });
3364 let write_task = SocketClientInner::spawn_write_task(
3365 Arc::clone(&connection_state),
3366 Arc::clone(&state_notify),
3367 initial_writer,
3368 writer_rx,
3369 b"\r\n".to_vec(),
3370 Some(sink),
3371 );
3372
3373 let (update_tx, update_rx) = oneshot::channel();
3374 let failing_writer: TestWriter = Box::pin(FailingWriter);
3375 writer_tx
3376 .send(WriterCommand::Update(failing_writer, update_tx))
3377 .unwrap();
3378 assert!(update_rx.await.unwrap());
3379
3380 writer_tx
3381 .send(WriterCommand::Send(Bytes::from_static(b"late")))
3382 .unwrap();
3383 yield_now().await;
3384 connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
3385 state_notify.notify_waiters();
3386 tokio::time::advance(Duration::from_millis(
3387 CONNECTION_STATE_CHECK_INTERVAL_MS * 2,
3388 ))
3389 .await;
3390 yield_now().await;
3391
3392 assert_eq!(
3393 ConnectionMode::from_atomic(&connection_state),
3394 ConnectionMode::Reconnect
3395 );
3396 assert_eq!(*states.lock(), vec![SocketState::Disconnected]);
3397
3398 let recorded = Arc::new(Mutex::new(Vec::new()));
3399 let replacement: TestWriter = Box::pin(RecordingWriter {
3400 bytes: Arc::clone(&recorded),
3401 });
3402 let (retry_tx, retry_rx) = oneshot::channel();
3403 writer_tx
3404 .send(WriterCommand::Update(replacement, retry_tx))
3405 .unwrap();
3406 assert!(retry_rx.await.unwrap());
3407
3408 connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
3409 state_notify.notify_waiters();
3410 drop(writer_tx);
3411 write_task.await.unwrap();
3412
3413 assert_eq!(recorded.lock().as_slice(), b"late\r\n");
3414 }
3415
3416 #[rstest]
3417 #[tokio::test(start_paused = true)]
3418 async fn test_failed_send_or_replay_is_not_duplicated_after_replay() {
3419 type TestWriter = Pin<Box<dyn AsyncWrite + Send>>;
3420
3421 let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
3422 let state_notify = Arc::new(tokio::sync::Notify::new());
3423 let (stream, _non_reading_peer) = tokio::io::duplex(1);
3424 let (pending_tx, pending_rx) = oneshot::channel();
3425 let writer: TestWriter = Box::pin(BackpressuredWriter {
3426 stream,
3427 pending_tx: Some(pending_tx),
3428 });
3429 let (writer_tx, writer_rx) =
3430 tokio::sync::mpsc::unbounded_channel::<WriterCommand<TestWriter>>();
3431 let write_task = SocketClientInner::spawn_write_task(
3432 Arc::clone(&connection_state),
3433 Arc::clone(&state_notify),
3434 writer,
3435 writer_rx,
3436 b"\r\n".to_vec(),
3437 None,
3438 );
3439 let subscription = Bytes::from_static(b"subscription");
3440
3441 writer_tx
3442 .send(WriterCommand::SendOrReplay {
3443 key: 7,
3444 data: subscription.clone(),
3445 })
3446 .unwrap();
3447 pending_rx.await.unwrap();
3448
3449 let recorded = Arc::new(Mutex::new(Vec::new()));
3450 let new_writer: TestWriter = Box::pin(RecordingWriter {
3451 bytes: Arc::clone(&recorded),
3452 });
3453 let (update_tx, update_rx) = oneshot::channel();
3454 writer_tx
3455 .send(WriterCommand::UpdateWithReplay(
3456 new_writer,
3457 vec![subscription],
3458 update_tx,
3459 ))
3460 .unwrap();
3461
3462 tokio::time::advance(Duration::from_secs(WRITE_TIMEOUT_SECS)).await;
3463 yield_now().await;
3464 assert!(update_rx.await.unwrap());
3465 assert_eq!(recorded.lock().as_slice(), b"subscription\r\n");
3466
3467 connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
3468 state_notify.notify_waiters();
3469 drop(writer_tx);
3470 write_task.await.unwrap();
3471 }
3472
3473 #[rstest]
3474 #[tokio::test]
3475 async fn test_connect_url_rejects_invalid_reconnect_backoff_before_connect() {
3476 let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
3477 listener.set_nonblocking(true).unwrap();
3478 let port = listener.local_addr().unwrap().port();
3479 let config = SocketConfig {
3480 url: format!("127.0.0.1:{port}"),
3481 mode: Mode::Plain,
3482 suffix: b"\r\n".to_vec(),
3483 message_handler: None,
3484 heartbeat: None,
3485 connect_timeout_ms: Some(1_000),
3486 reconnect_delay_initial_ms: Some(50),
3487 reconnect_delay_max_ms: Some(100),
3488 reconnect_backoff_factor: Some(100.1),
3489 reconnect_jitter_ms: Some(0),
3490 connection_max_retries: Some(1),
3491 reconnect_max_attempts: None,
3492 heartbeat_timeout_secs: None,
3493 certs_dir: None,
3494 };
3495
3496 let error = match SocketClientInner::connect_url(config, None).await {
3497 Ok(_) => panic!("invalid reconnect backoff should be rejected"),
3498 Err(e) => e,
3499 };
3500
3501 assert!(
3502 error.to_string().contains("factor"),
3503 "error should mention the invalid factor, was: {error}"
3504 );
3505 assert_eq!(
3506 listener.accept().unwrap_err().kind(),
3507 std::io::ErrorKind::WouldBlock,
3508 "invalid reconnect backoff must be rejected before connecting"
3509 );
3510 }
3511
3512 #[rstest]
3513 #[tokio::test]
3514 async fn test_reconnect_then_close() {
3515 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3517 let port = listener.local_addr().unwrap().port();
3518
3519 let server = task::spawn(async move {
3521 if let Ok((mut sock, _)) = listener.accept().await {
3522 drop(sock.shutdown());
3523 }
3524 sleep(Duration::from_secs(1)).await;
3526 });
3527
3528 let config = SocketConfig {
3530 url: format!("127.0.0.1:{port}"),
3531 mode: Mode::Plain,
3532 suffix: b"\r\n".to_vec(),
3533 message_handler: None,
3534 heartbeat: None,
3535 connect_timeout_ms: Some(1_000),
3536 reconnect_delay_initial_ms: Some(50),
3537 reconnect_delay_max_ms: Some(100),
3538 reconnect_backoff_factor: Some(1.0),
3539 reconnect_jitter_ms: Some(0),
3540 connection_max_retries: Some(1),
3541 reconnect_max_attempts: None,
3542 heartbeat_timeout_secs: None,
3543 certs_dir: None,
3544 };
3545
3546 let client = SocketClient::builder()
3548 .config(config.clone())
3549 .connect()
3550 .await
3551 .unwrap();
3552
3553 wait_until_async(
3555 || async { client.is_reconnecting() },
3556 Duration::from_secs(2),
3557 )
3558 .await;
3559
3560 client.close().await;
3562 assert!(client.is_closed());
3563 server.abort();
3564 }
3565
3566 #[rstest]
3567 #[tokio::test]
3568 async fn test_reconnect_state_flips_when_reader_stops() {
3569 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3571 let port = listener.local_addr().unwrap().port();
3572
3573 let server = task::spawn(async move {
3574 if let Ok((sock, _)) = listener.accept().await {
3575 drop(sock);
3576 }
3577 sleep(Duration::from_millis(50)).await;
3579 });
3580
3581 let config = SocketConfig {
3582 url: format!("127.0.0.1:{port}"),
3583 mode: Mode::Plain,
3584 suffix: b"\r\n".to_vec(),
3585 message_handler: None,
3586 heartbeat: None,
3587 connect_timeout_ms: Some(1_000),
3588 reconnect_delay_initial_ms: Some(50),
3589 reconnect_delay_max_ms: Some(100),
3590 reconnect_backoff_factor: Some(1.0),
3591 reconnect_jitter_ms: Some(0),
3592 connection_max_retries: Some(1),
3593 reconnect_max_attempts: None,
3594 heartbeat_timeout_secs: None,
3595 certs_dir: None,
3596 };
3597
3598 let client = SocketClient::builder()
3599 .config(config)
3600 .connect()
3601 .await
3602 .unwrap();
3603
3604 wait_until_async(
3605 || async { client.is_reconnecting() },
3606 Duration::from_secs(2),
3607 )
3608 .await;
3609
3610 client.close().await;
3611 server.abort();
3612 }
3613
3614 #[rstest]
3615 fn test_parse_socket_url_raw_address() {
3616 let (socket_addr, request_url) =
3618 SocketClientInner::parse_socket_url("example.com:6130", Mode::Tls).unwrap();
3619 assert_eq!(socket_addr, "example.com:6130");
3620 assert_eq!(request_url, "wss://example.com:6130");
3621
3622 let (socket_addr, request_url) =
3624 SocketClientInner::parse_socket_url("localhost:8080", Mode::Plain).unwrap();
3625 assert_eq!(socket_addr, "localhost:8080");
3626 assert_eq!(request_url, "ws://localhost:8080");
3627 }
3628
3629 #[rstest]
3630 fn test_parse_socket_url_with_scheme() {
3631 let (socket_addr, request_url) =
3633 SocketClientInner::parse_socket_url("wss://example.com:443/path", Mode::Tls).unwrap();
3634 assert_eq!(socket_addr, "example.com:443");
3635 assert_eq!(request_url, "wss://example.com:443/path");
3636
3637 let (socket_addr, request_url) =
3639 SocketClientInner::parse_socket_url("ws://localhost:8080", Mode::Plain).unwrap();
3640 assert_eq!(socket_addr, "localhost:8080");
3641 assert_eq!(request_url, "ws://localhost:8080");
3642 }
3643
3644 #[rstest]
3645 fn test_parse_socket_url_default_ports() {
3646 let (socket_addr, _) =
3648 SocketClientInner::parse_socket_url("wss://example.com", Mode::Tls).unwrap();
3649 assert_eq!(socket_addr, "example.com:443");
3650
3651 let (socket_addr, _) =
3653 SocketClientInner::parse_socket_url("ws://example.com", Mode::Plain).unwrap();
3654 assert_eq!(socket_addr, "example.com:80");
3655
3656 let (socket_addr, _) =
3658 SocketClientInner::parse_socket_url("https://example.com", Mode::Tls).unwrap();
3659 assert_eq!(socket_addr, "example.com:443");
3660
3661 let (socket_addr, _) =
3663 SocketClientInner::parse_socket_url("http://example.com", Mode::Plain).unwrap();
3664 assert_eq!(socket_addr, "example.com:80");
3665 }
3666
3667 #[rstest]
3668 fn test_parse_socket_url_unknown_scheme_uses_mode() {
3669 let (socket_addr, _) =
3671 SocketClientInner::parse_socket_url("custom://example.com", Mode::Tls).unwrap();
3672 assert_eq!(socket_addr, "example.com:443");
3673
3674 let (socket_addr, _) =
3675 SocketClientInner::parse_socket_url("custom://example.com", Mode::Plain).unwrap();
3676 assert_eq!(socket_addr, "example.com:80");
3677 }
3678
3679 #[rstest]
3680 fn test_parse_socket_url_ipv6() {
3681 let (socket_addr, request_url) =
3683 SocketClientInner::parse_socket_url("[::1]:8080", Mode::Plain).unwrap();
3684 assert_eq!(socket_addr, "[::1]:8080");
3685 assert_eq!(request_url, "ws://[::1]:8080");
3686
3687 let (socket_addr, _) =
3689 SocketClientInner::parse_socket_url("ws://[::1]:8080", Mode::Plain).unwrap();
3690 assert_eq!(socket_addr, "[::1]:8080");
3691 }
3692
3693 #[rstest]
3694 #[tokio::test]
3695 async fn test_url_parsing_raw_socket_address() {
3696 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3698 let port = listener.local_addr().unwrap().port();
3699
3700 let server = task::spawn(async move {
3701 if let Ok((sock, _)) = listener.accept().await {
3702 drop(sock);
3703 }
3704 sleep(Duration::from_millis(50)).await;
3705 });
3706
3707 let config = SocketConfig {
3708 url: format!("127.0.0.1:{port}"), mode: Mode::Plain,
3710 suffix: b"\r\n".to_vec(),
3711 message_handler: None,
3712 heartbeat: None,
3713 connect_timeout_ms: Some(1_000),
3714 reconnect_delay_initial_ms: Some(50),
3715 reconnect_delay_max_ms: Some(100),
3716 reconnect_backoff_factor: Some(1.0),
3717 reconnect_jitter_ms: Some(0),
3718 connection_max_retries: Some(1),
3719 reconnect_max_attempts: None,
3720 heartbeat_timeout_secs: None,
3721 certs_dir: None,
3722 };
3723
3724 let client = SocketClient::builder().config(config).connect().await;
3726 assert!(
3727 client.is_ok(),
3728 "Client should connect with raw socket address format"
3729 );
3730
3731 if let Ok(client) = client {
3732 client.close().await;
3733 }
3734 server.abort();
3735 }
3736
3737 #[rstest]
3738 #[tokio::test]
3739 async fn test_url_parsing_with_scheme() {
3740 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3742 let port = listener.local_addr().unwrap().port();
3743
3744 let server = task::spawn(async move {
3745 if let Ok((sock, _)) = listener.accept().await {
3746 drop(sock);
3747 }
3748 sleep(Duration::from_millis(50)).await;
3749 });
3750
3751 let config = SocketConfig {
3752 url: format!("ws://127.0.0.1:{port}"), mode: Mode::Plain,
3754 suffix: b"\r\n".to_vec(),
3755 message_handler: None,
3756 heartbeat: None,
3757 connect_timeout_ms: Some(1_000),
3758 reconnect_delay_initial_ms: Some(50),
3759 reconnect_delay_max_ms: Some(100),
3760 reconnect_backoff_factor: Some(1.0),
3761 reconnect_jitter_ms: Some(0),
3762 connection_max_retries: Some(1),
3763 reconnect_max_attempts: None,
3764 heartbeat_timeout_secs: None,
3765 certs_dir: None,
3766 };
3767
3768 let client = SocketClient::builder().config(config).connect().await;
3770 assert!(
3771 client.is_ok(),
3772 "Client should connect with URL scheme format"
3773 );
3774
3775 if let Ok(client) = client {
3776 client.close().await;
3777 }
3778 server.abort();
3779 }
3780
3781 #[rstest]
3782 fn test_parse_socket_url_ipv6_with_zone() {
3783 let (socket_addr, request_url) =
3785 SocketClientInner::parse_socket_url("[fe80::1%eth0]:8080", Mode::Plain).unwrap();
3786 assert_eq!(socket_addr, "[fe80::1%eth0]:8080");
3787 assert_eq!(request_url, "ws://[fe80::1%eth0]:8080");
3788
3789 let (socket_addr, request_url) =
3791 SocketClientInner::parse_socket_url("ws://[fe80::1%lo]:9090", Mode::Plain).unwrap();
3792 assert_eq!(socket_addr, "[fe80::1%lo]:9090");
3793 assert_eq!(request_url, "ws://[fe80::1%lo]:9090");
3794 }
3795
3796 #[rstest]
3797 #[tokio::test]
3798 async fn test_ipv6_loopback_connection() {
3799 if TcpListener::bind("[::1]:0").await.is_err() {
3802 return;
3803 }
3804
3805 let listener = TcpListener::bind("[::1]:0").await.unwrap();
3806 let port = listener.local_addr().unwrap().port();
3807
3808 let server = task::spawn(async move {
3809 if let Ok((mut sock, _)) = listener.accept().await {
3810 let mut buf = vec![0u8; 1024];
3811 if let Ok(n) = sock.read(&mut buf).await {
3812 let _ = sock.write_all(&buf[..n]).await;
3814 }
3815 }
3816 sleep(Duration::from_millis(50)).await;
3817 });
3818
3819 let config = SocketConfig {
3820 url: format!("[::1]:{port}"), mode: Mode::Plain,
3822 suffix: b"\r\n".to_vec(),
3823 message_handler: None,
3824 heartbeat: None,
3825 connect_timeout_ms: Some(1_000),
3826 reconnect_delay_initial_ms: Some(50),
3827 reconnect_delay_max_ms: Some(100),
3828 reconnect_backoff_factor: Some(1.0),
3829 reconnect_jitter_ms: Some(0),
3830 connection_max_retries: Some(1),
3831 reconnect_max_attempts: None,
3832 heartbeat_timeout_secs: None,
3833 certs_dir: None,
3834 };
3835
3836 let client = SocketClient::builder().config(config).connect().await;
3837 assert!(
3838 client.is_ok(),
3839 "Client should connect to IPv6 loopback address"
3840 );
3841
3842 if let Ok(client) = client {
3843 client.close().await;
3844 }
3845 server.abort();
3846 }
3847
3848 #[rstest]
3849 #[tokio::test]
3850 async fn test_send_waits_during_reconnection() {
3851 use nautilus_common::testing::wait_until_async;
3853
3854 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3855 let port = listener.local_addr().unwrap().port();
3856
3857 let server = task::spawn(async move {
3858 if let Ok((sock, _)) = listener.accept().await {
3860 drop(sock);
3861 }
3862
3863 sleep(Duration::from_millis(500)).await;
3865
3866 if let Ok((mut sock, _)) = listener.accept().await {
3868 let mut buf = vec![0u8; 1024];
3870 while let Ok(n) = sock.read(&mut buf).await {
3871 if n == 0 {
3872 break;
3873 }
3874
3875 if sock.write_all(&buf[..n]).await.is_err() {
3876 break;
3877 }
3878 }
3879 }
3880 });
3881
3882 let config = SocketConfig {
3883 url: format!("127.0.0.1:{port}"),
3884 mode: Mode::Plain,
3885 suffix: b"\r\n".to_vec(),
3886 message_handler: None,
3887 heartbeat: None,
3888 connect_timeout_ms: Some(5_000), reconnect_delay_initial_ms: Some(100),
3890 reconnect_delay_max_ms: Some(200),
3891 reconnect_backoff_factor: Some(1.0),
3892 reconnect_jitter_ms: Some(0),
3893 connection_max_retries: Some(1),
3894 reconnect_max_attempts: None,
3895 heartbeat_timeout_secs: None,
3896 certs_dir: None,
3897 };
3898
3899 let client = SocketClient::builder()
3900 .config(config)
3901 .connect()
3902 .await
3903 .unwrap();
3904
3905 wait_until_async(
3907 || async { client.is_reconnecting() },
3908 Duration::from_secs(2),
3909 )
3910 .await;
3911
3912 let send_result = tokio::time::timeout(
3914 Duration::from_secs(3),
3915 client.send_bytes(b"test_message".to_vec()),
3916 )
3917 .await;
3918
3919 assert!(
3920 send_result.is_ok() && send_result.unwrap().is_ok(),
3921 "Send should succeed after waiting for reconnection"
3922 );
3923
3924 client.close().await;
3925 server.abort();
3926 }
3927
3928 #[rstest]
3929 #[tokio::test]
3930 async fn test_send_bytes_timeout_uses_configured_connect_timeout() {
3931 use nautilus_common::testing::wait_until_async;
3934
3935 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3936 let port = listener.local_addr().unwrap().port();
3937
3938 let server = task::spawn(async move {
3939 if let Ok((sock, _)) = listener.accept().await {
3941 drop(sock);
3942 }
3943 drop(listener);
3945 sleep(Duration::from_mins(1)).await;
3946 });
3947
3948 let config = SocketConfig {
3949 url: format!("127.0.0.1:{port}"),
3950 mode: Mode::Plain,
3951 suffix: b"\r\n".to_vec(),
3952 message_handler: None,
3953 heartbeat: None,
3954 connect_timeout_ms: Some(1_000), reconnect_delay_initial_ms: Some(200), reconnect_delay_max_ms: Some(200),
3957 reconnect_backoff_factor: Some(1.0),
3958 reconnect_jitter_ms: Some(0),
3959 connection_max_retries: Some(1),
3960 reconnect_max_attempts: None,
3961 heartbeat_timeout_secs: None,
3962 certs_dir: None,
3963 };
3964
3965 let client = SocketClient::builder()
3966 .config(config)
3967 .connect()
3968 .await
3969 .unwrap();
3970
3971 wait_until_async(
3973 || async { client.is_reconnecting() },
3974 Duration::from_secs(3),
3975 )
3976 .await;
3977
3978 let start = std::time::Instant::now();
3981 let send_result = client.send_bytes(b"test".to_vec()).await;
3982 let elapsed = start.elapsed();
3983
3984 assert!(
3985 send_result.is_err(),
3986 "Send should fail when client stuck in RECONNECT, was: {send_result:?}"
3987 );
3988 assert!(
3989 matches!(send_result, Err(crate::error::SendError::Timeout)),
3990 "Send should return Timeout error, was: {send_result:?}"
3991 );
3992 assert!(
3995 elapsed >= Duration::from_millis(900),
3996 "Send should timeout after at least 1s (configured timeout), took {elapsed:?}"
3997 );
3998
3999 client.close().await;
4000 server.abort();
4001 }
4002
4003 #[rstest]
4004 #[tokio::test]
4005 async fn test_heartbeat_timeout_triggers_reconnect() {
4006 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4007 let port = listener.local_addr().unwrap().port();
4008
4009 let server = task::spawn(async move {
4011 let (_sock1, _) = listener.accept().await.unwrap();
4012 sleep(Duration::from_secs(5)).await;
4014 });
4015
4016 let config = SocketConfig {
4017 url: format!("127.0.0.1:{port}"),
4018 mode: Mode::Plain,
4019 suffix: b"\r\n".to_vec(),
4020 message_handler: None,
4021 heartbeat: None,
4022 connect_timeout_ms: Some(2_000),
4023 reconnect_delay_initial_ms: Some(50),
4024 reconnect_delay_max_ms: Some(100),
4025 reconnect_backoff_factor: Some(1.0),
4026 reconnect_jitter_ms: Some(0),
4027 connection_max_retries: Some(1),
4028 reconnect_max_attempts: Some(1),
4029 heartbeat_timeout_secs: Some(1),
4030 certs_dir: None,
4031 };
4032
4033 let client = SocketClient::builder()
4034 .config(config)
4035 .connect()
4036 .await
4037 .unwrap();
4038
4039 assert!(client.is_active());
4040
4041 wait_until_async(
4043 || async { client.is_reconnecting() || client.is_closed() },
4044 Duration::from_secs(4),
4045 )
4046 .await;
4047
4048 assert!(
4049 !client.is_active(),
4050 "Client should not be active after the dead-peer timeout"
4051 );
4052
4053 client.close().await;
4054 server.abort();
4055 }
4056
4057 #[rstest]
4058 #[tokio::test]
4059 async fn test_heartbeat_timeout_resets_on_inbound_bytes() {
4060 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4061 let port = listener.local_addr().unwrap().port();
4062
4063 let server = task::spawn(async move {
4065 let (mut sock, _) = listener.accept().await.unwrap();
4066 for _ in 0..10 {
4067 sleep(Duration::from_millis(200)).await;
4068
4069 if sock.write_all(b"ping\r\n").await.is_err() {
4070 break;
4071 }
4072 }
4073 });
4074
4075 let config = SocketConfig {
4076 url: format!("127.0.0.1:{port}"),
4077 mode: Mode::Plain,
4078 suffix: b"\r\n".to_vec(),
4079 message_handler: None,
4080 heartbeat: None,
4081 connect_timeout_ms: Some(2_000),
4082 reconnect_delay_initial_ms: Some(50),
4083 reconnect_delay_max_ms: Some(100),
4084 reconnect_backoff_factor: Some(1.0),
4085 reconnect_jitter_ms: Some(0),
4086 connection_max_retries: Some(1),
4087 reconnect_max_attempts: Some(1),
4088 heartbeat_timeout_secs: Some(1),
4089 certs_dir: None,
4090 };
4091
4092 let client = SocketClient::builder()
4093 .config(config)
4094 .connect()
4095 .await
4096 .unwrap();
4097
4098 assert!(client.is_active());
4099
4100 sleep(Duration::from_millis(1_500)).await;
4102
4103 assert!(
4104 client.is_active(),
4105 "Client should remain active when data is flowing"
4106 );
4107
4108 client.close().await;
4109 server.abort();
4110 }
4111
4112 #[rstest]
4113 #[tokio::test]
4114 async fn test_close_during_backoff_exits_promptly() {
4115 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4119 let port = listener.local_addr().unwrap().port();
4120
4121 let server = task::spawn(async move {
4122 if let Ok((mut sock, _)) = listener.accept().await {
4124 drop(sock.shutdown());
4125 }
4126 sleep(Duration::from_mins(1)).await;
4128 });
4129
4130 let config = SocketConfig {
4131 url: format!("127.0.0.1:{port}"),
4132 mode: Mode::Plain,
4133 suffix: b"\r\n".to_vec(),
4134 message_handler: None,
4135 heartbeat: None,
4136 connect_timeout_ms: Some(1_000),
4137 reconnect_delay_initial_ms: Some(10_000), reconnect_delay_max_ms: Some(10_000),
4139 reconnect_backoff_factor: Some(1.0),
4140 reconnect_jitter_ms: Some(0),
4141 connection_max_retries: None,
4142 reconnect_max_attempts: None,
4143 heartbeat_timeout_secs: None,
4144 certs_dir: None,
4145 };
4146
4147 let client = SocketClient::builder()
4148 .config(config)
4149 .connect()
4150 .await
4151 .unwrap();
4152
4153 wait_until_async(
4155 || async { client.is_reconnecting() },
4156 Duration::from_secs(3),
4157 )
4158 .await;
4159
4160 sleep(Duration::from_millis(1_500)).await;
4162
4163 let start = std::time::Instant::now();
4165 client.close().await;
4166 let elapsed = start.elapsed();
4167
4168 assert!(client.is_closed(), "Client should be closed");
4169 assert!(
4171 elapsed < Duration::from_secs(2),
4172 "Close should interrupt backoff sleep, took {elapsed:?}"
4173 );
4174
4175 server.abort();
4176 }
4177
4178 #[rstest]
4179 #[tokio::test]
4180 async fn test_zero_heartbeat_timeout_rejected() {
4181 let config = SocketConfig {
4182 url: "127.0.0.1:9999".to_string(),
4183 mode: Mode::Plain,
4184 suffix: b"\r\n".to_vec(),
4185 message_handler: None,
4186 heartbeat: None,
4187 connect_timeout_ms: None,
4188 reconnect_delay_initial_ms: None,
4189 reconnect_delay_max_ms: None,
4190 reconnect_backoff_factor: None,
4191 reconnect_jitter_ms: None,
4192 reconnect_max_attempts: None,
4193 connection_max_retries: Some(1),
4194 heartbeat_timeout_secs: Some(0),
4195 certs_dir: None,
4196 };
4197
4198 let result = SocketClient::builder().config(config).connect().await;
4199
4200 assert!(result.is_err(), "Zero heartbeat timeout should be rejected");
4201 let err_msg = result.unwrap_err().to_string();
4202 assert!(
4203 err_msg.contains("heartbeat_timeout_secs"),
4204 "Error should name the offending field, was: {err_msg}"
4205 );
4206 }
4207
4208 #[rstest]
4209 #[tokio::test]
4210 async fn test_empty_suffix_rejected() {
4211 let config = SocketConfig {
4212 url: "127.0.0.1:9999".to_string(),
4213 mode: Mode::Plain,
4214 suffix: vec![],
4215 message_handler: None,
4216 heartbeat: None,
4217 connect_timeout_ms: None,
4218 reconnect_delay_initial_ms: None,
4219 reconnect_delay_max_ms: None,
4220 reconnect_backoff_factor: None,
4221 reconnect_jitter_ms: None,
4222 reconnect_max_attempts: None,
4223 connection_max_retries: Some(1),
4224 heartbeat_timeout_secs: None,
4225 certs_dir: None,
4226 };
4227
4228 let result = SocketClient::builder().config(config).connect().await;
4229
4230 assert!(
4231 result.is_err(),
4232 "Empty suffix should cause connection to fail"
4233 );
4234 let err_msg = result.unwrap_err().to_string();
4235 assert!(
4236 err_msg.contains("suffix cannot be empty"),
4237 "Error should mention empty suffix, was: {err_msg}"
4238 );
4239 }
4240}
4241
4242#[cfg(test)]
4243mod reconnect_request_tests {
4244 use std::sync::{Arc, atomic::AtomicU8};
4245
4246 use parking_lot::Mutex;
4247 use rstest::rstest;
4248
4249 use super::*;
4250 use crate::SocketState;
4251
4252 fn handle(
4253 mode: ConnectionMode,
4254 ) -> (
4255 SocketReconnectHandle,
4256 Arc<tokio::sync::Notify>,
4257 Arc<Mutex<Vec<SocketState>>>,
4258 ) {
4259 let controller_notify = Arc::new(tokio::sync::Notify::new());
4260 let states = Arc::new(Mutex::new(Vec::new()));
4261 let states_callback = Arc::clone(&states);
4262 let state_sink = SocketStateSink::new(move |state| {
4263 states_callback.lock().push(state);
4264 });
4265 let handle = SocketReconnectHandle {
4266 connection_mode: Arc::new(AtomicU8::new(mode.as_u8())),
4267 state_sink: Some(state_sink),
4268 controller_lifecycle: Arc::new(ControllerLifecycle::new()),
4269 controller_notify: Arc::clone(&controller_notify),
4270 };
4271 (handle, controller_notify, states)
4272 }
4273
4274 #[rstest]
4275 #[tokio::test]
4276 async fn accepted_request_reports_loss_and_wakes_controller_once() {
4277 let (handle, controller_notify, states) = handle(ConnectionMode::Active);
4278
4279 assert_eq!(
4280 handle.request_reconnect(),
4281 ReconnectRequestOutcome::Accepted
4282 );
4283 assert_eq!(*states.lock(), vec![SocketState::Disconnected]);
4284 tokio::time::timeout(Duration::from_millis(10), controller_notify.notified())
4285 .await
4286 .expect("accepted request should notify controller");
4287
4288 assert_eq!(
4289 handle.request_reconnect(),
4290 ReconnectRequestOutcome::AlreadyReconnecting
4291 );
4292 assert_eq!(*states.lock(), vec![SocketState::Disconnected]);
4293 assert!(
4294 tokio::time::timeout(Duration::from_millis(10), controller_notify.notified())
4295 .await
4296 .is_err(),
4297 "duplicate request should not notify controller",
4298 );
4299 }
4300
4301 #[rstest]
4302 #[case(
4303 ConnectionMode::Reconnect,
4304 ReconnectRequestOutcome::AlreadyReconnecting
4305 )]
4306 #[case(ConnectionMode::Disconnect, ReconnectRequestOutcome::Disconnected)]
4307 #[case(ConnectionMode::Closed, ReconnectRequestOutcome::Closed)]
4308 #[tokio::test]
4309 async fn rejected_request_preserves_state_and_does_not_wake_controller(
4310 #[case] mode: ConnectionMode,
4311 #[case] expected: ReconnectRequestOutcome,
4312 ) {
4313 let (handle, controller_notify, states) = handle(mode);
4314
4315 assert_eq!(handle.request_reconnect(), expected);
4316 assert!(states.lock().is_empty());
4317 assert!(
4318 tokio::time::timeout(Duration::from_millis(10), controller_notify.notified())
4319 .await
4320 .is_err(),
4321 "rejected request should not notify controller",
4322 );
4323 }
4324}