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